多模态 AI 能够处理和理解多种类型的数据:
import torch import torch.nn as nn from transformers import CLIPModel, CLIPTokenizer, CLIPProcessor # 加载预训练模型 model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32") processor = CLIPProcessor.from_pretrained("openai/clip-vit-base-patch32") # 编码文本 text_inputs = processor(text=["一只狗", "一只猫", "一只鸟"], return_tensors="pt", padding=True) text_features = model.get_text_features(**text_inputs) # 编码图像 image_inputs = processor(images=[dog_image, cat_image, bird_image], return_tensors="pt") image_features = model.get_image_features(**image_inputs) # 计算相似度 logits_per_image = model.logit_per_image(text_features, image_features) probs = torch.softmax(logits_per_image, dim=-1) print(f"文本-图像匹配概率:{probs}")
class CLIPModel(nn.Module): def __init__(self): super().__init__() # 图像编码器(Vision Encoder) self.vision_encoder = VisionEncoder() # 文本编码器(Text Encoder) self.text_encoder = TextEncoder() # 对比学习 self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07)) def forward(self, images, text): # 编码 image_features = self.vision_encoder(images) text_features = self.text_encoder(text) # 归一化 image_features = image_features / image_features.norm(dim=1, keepdim=True) text_features = text_features / text_features.norm(dim=1, keepdim=True) # 计算相似度 logits = torch.matmul(text_features, image_features.T) * self.logit_scale return logits
class ContrastiveLoss(nn.Module): def __init__(self, temperature=0.07): super().__init__() self.temperature = temperature def forward(self, image_features, text_features): # 计算相似度矩阵 logits = torch.matmul(image_features, text_features.T) / self.temperature # 标签:对角线为正样本 batch_size = image_features.size(0) labels = torch.arange(batch_size) # 交叉熵损失 loss_i = F.cross_entropy(logits, labels) loss_t = F.cross_entropy(logits.T, labels) loss = (loss_i + loss_t) / 2 return loss
from torchvision import transforms transform = transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.4), transforms.RandomRotation(15), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])
class ImageSearchEngine: def __init__(self, clip_model, vector_store): self.clip_model = clip_model self.vector_store = vector_store def index_images(self, images, descriptions): """索引图像""" # 编码图像 image_features = self.clip_model.encode_images(images) # 存储到向量数据库 for feature, desc in zip(image_features, descriptions): self.vector_store.add(feature, { 'type': 'image', 'description': desc }) def search(self, query_text, top_k=10): """文本搜图""" # 编码查询文本 query_feature = self.clip_model.encode_text(query_text) # 向量搜索 results = self.vector_store.search(query_feature, top_k=top_k) return results
class ZeroShotClassifier: def __init__(self, clip_model, class_names): self.clip_model = clip_model self.class_names = class_names self.class_features = self.clip_model.encode_text(class_names) def classify(self, image): """零样本分类""" # 编码图像 image_feature = self.clip_model.encode_image(image) # 计算与各类别的相似度 similarities = F.cosine_similarity( image_feature.unsqueeze(0), self.class_features ) # 返回最相似的类别 class_idx = similarities.argmax() return self.class_names[class_idx], similarities[class_idx]
from transformers import BlipForImageTextGeneration from PIL import Image # 加载模型 model = BlipForImageTextGeneration.from_pretrained("Salesforce/blip-image-captioning-base") processor = BlipProcessor.from_pretrained("Salesforce/blip-image-captioning-base") # 生成描述 def generate_caption(image_path): image = Image.open(image_path) inputs = processor(image, return_tensors="pt") with torch.no_grad(): outputs = model.generate(**inputs) caption = processor.decode(outputs[0]) return caption # 使用示例 caption = generate_caption("path/to/image.jpg") print(f"图像描述:{caption}")
# 量化模型以减少内存和加速推理 from transformers import BitsAndBytesConfig quantization_config = BitsAndBytesConfig( load_in_8bit=True, llm_int8_threshold=6.0 ) model = CLIPModel.from_pretrained( "openai/clip-vit-base-patch32", quantization_config=quantization_config )
def batch_inference(model, images, batch_size=32): """批量推理""" results = [] for i in range(0, len(images), batch_size): batch = images[i:i+batch_size] # 批量编码 features = model.encode_images(batch) results.extend(features) return results
class EarlyFusion(nn.Module): def __init__(self): super().__init__() self.image_encoder = VisionEncoder() self.text_encoder = TextEncoder() self.fusion_layer = FusionLayer() def forward(self, image, text): # 编码 image_feat = self.image_encoder(image) text_feat = self.text_encoder(text) # 早期融合:拼接后再处理 combined = torch.cat([image_feat, text_feat], dim=-1) output = self.fusion_layer(combined) return output
class LateFusion(nn.Module): def __init__(self): super().__init__() self.image_encoder = visions_encoder self.text_encoder = text_encoder self.classifier = nn.Linear(512, num_classes) def forward(self, image, text): # 分别编码 image_feat = self.image_encoder(image) text_feat = self.text_encoder(text) # 后期融合:各自处理后融合 image_output = self.image_classifier(image_feat) text_output = self.text_classifier(text_feat) # 融合 combined = (image_output + text_output) / 2 return combined
class ProductSearch: def __init__(self, clip_model, product_database): self.clip_model = clip_model self.product_db = product_database # 索引商品图像 self.index_products() def index_products(self): products = self.product_db.get_all_products() images = [p['image'] for p in products] descriptions = [p['description'] for p in products] self.search_engine.index_images(images, descriptions) def search(self, query): results = self.search_engine.search(query, top_k=10) return [ self.product_db.get_product(result['id']) for result in results ]
class ContentModerator: def __"" 检测内容是否包含违规内容 """ def __init__(self, clip_model, forbidden_concepts): self.clip_model = clip_model self.forbidden_concepts = forbidden_concepts self.concept_features = self.clip_model.encode_text(forbidden_concepts) def check_image(self, image): """检查图像""" image_feature = self.clip_model.encode_image(image) # 计算与违规概念的相似度 similarities = F.cosine_similarity( image_feature.unsqueeze(0), self.concept_features ) max_similarity = similarities.max() if max_similarity > 0.8: return { 'safe': False, 'reason': f"检测到违规内容:{self.forbidden_concepts[similarities.argmax()]}", 'confidence': max_similarity } return {'safe': True}
多模态 AI 的关键:
掌握多模态 AI,构建更智能的应用!