import torch import torch.nn as nn import torch.nn.functional as F from torchvision import models, transforms from PIL import Image import numpy as np from typing import Union, List, Tuple, Optional from settings import settings class ResNet50EmbeddingNet(nn.Module): """ 基于ResNet50的食物特征提取网络 移除分类层,输出512维特征向量用于相似度计算 """ def __init__(self, embedding_dim: int = 512, pretrained: bool = True, use_internal_preprocess: bool = False): """ 初始化ResNet50 Embedding网络 Args: embedding_dim: 输出特征向量维度,默认512 pretrained: 是否使用预训练权重 use_internal_preprocess: 是否在模型内部进行预处理 """ super(ResNet50EmbeddingNet, self).__init__() self.embedding_dim = embedding_dim self.use_internal_preprocess = use_internal_preprocess # 加载预训练的ResNet50 self.backbone = models.resnet50(pretrained=pretrained) # 获取ResNet50最后一层的输入特征数(2048) backbone_output_dim = self.backbone.fc.in_features # 移除原始的分类层 self.backbone.fc = nn.Identity() # 添加embedding层 self.embedding_layer = nn.Sequential( nn.Linear(backbone_output_dim, embedding_dim), nn.BatchNorm1d(embedding_dim), nn.ReLU(inplace=True), nn.Dropout(0.2) ) # 图片预处理变换(用于推理) self.preprocess = transforms.Compose([ transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BILINEAR), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 内部预处理变换(用于已经是tensor但未归一化的数据) self.internal_preprocess = transforms.Compose([ transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) def preprocess_image(self, image: Union[Image.Image, np.ndarray]) -> torch.Tensor: """ 预处理单张图片 Args: image: PIL Image 或 numpy array Returns: torch.Tensor: 预处理后的张量,形状为 (1, 3, 224, 224) """ if not isinstance(image, Image.Image): # 如果是numpy array,转换为PIL Image if hasattr(image, 'shape'): if len(image.shape) == 3 and image.shape[2] == 3: image = Image.fromarray(image.astype(np.uint8)) else: raise ValueError("numpy array必须是3通道RGB格式") else: raise ValueError("输入必须是PIL Image或numpy array") # 应用预处理变换 tensor = self.preprocess(image) # 添加batch维度 tensor = tensor.unsqueeze(0) return tensor def mobile_preprocess(self, x: torch.Tensor) -> torch.Tensor: """ 移动端预处理(TorchScript兼容) Args: x: 输入tensor,形状为 [batch, 3, height, width],值范围 0-1 Returns: torch.Tensor: 预处理后的tensor,形状为 [batch, 3, 224, 224] """ # 缩放到 224x224 x = F.interpolate(x, size=(224, 224), mode='bilinear', align_corners=False) # ImageNet标准化 mean = torch.tensor([0.485, 0.456, 0.406], device=x.device).view(1, 3, 1, 1) std = torch.tensor([0.229, 0.224, 0.225], device=x.device).view(1, 3, 1, 1) x = (x - mean) / std return x def _forward_backbone(self, x: torch.Tensor) -> torch.Tensor: """ 通过ResNet50骨干网络提取特征 Args: x: 预处理后的输入tensor Returns: torch.Tensor: 骨干网络输出特征 """ return self.backbone(x) def _forward_embedding(self, backbone_features: torch.Tensor) -> torch.Tensor: """ 将骨干网络特征转换为embedding向量 Args: backbone_features: 骨干网络输出特征 Returns: torch.Tensor: embedding向量 """ return self.embedding_layer(backbone_features) def forward(self, x: torch.Tensor, normalize: bool = True) -> torch.Tensor: """ 前向传播 Args: x: 输入tensor normalize: 是否对输出进行L2归一化 Returns: torch.Tensor: embedding向量 """ # 如果启用内部预处理且输入是tensor if self.use_internal_preprocess and isinstance(x, torch.Tensor): x = self.internal_preprocess(x) # 通过骨干网络 backbone_features = self._forward_backbone(x) # 生成embedding embeddings = self._forward_embedding(backbone_features) # L2归一化(用于余弦相似度计算) if normalize: embeddings = F.normalize(embeddings, p=2, dim=1) return embeddings def forward_mobile(self, x: torch.Tensor, normalize: bool = True) -> torch.Tensor: """ 移动端前向传播(包含预处理) Args: x: 输入tensor,形状为 [batch, 3, height, width],值范围 0-1 normalize: 是否对输出进行L2归一化 Returns: torch.Tensor: embedding向量 """ # 移动端预处理 x = self.mobile_preprocess(x) # 执行前向传播 return self.forward(x, normalize=normalize) def extract_embedding(self, image: Union[Image.Image, np.ndarray, torch.Tensor], normalize: bool = True) -> np.ndarray: """ 从单张图片提取embedding向量 Args: image: 输入图片(PIL Image、numpy array或torch.Tensor) normalize: 是否进行L2归一化 Returns: np.ndarray: embedding向量 """ self.eval() with torch.no_grad(): if isinstance(image, torch.Tensor): if len(image.shape) == 3: image = image.unsqueeze(0) # 添加batch维度 x = image else: x = self.preprocess_image(image) # 移动到正确的设备 device = next(self.parameters()).device x = x.to(device) # 提取特征 embedding = self.forward(x, normalize=normalize) return embedding.cpu().numpy().flatten() def extract_batch_embeddings(self, images: List[Union[Image.Image, np.ndarray]], normalize: bool = True, batch_size: int = 32) -> np.ndarray: """ 批量提取embedding向量 Args: images: 图片列表 normalize: 是否进行L2归一化 batch_size: 批处理大小 Returns: np.ndarray: embedding向量数组,形状为 (N, embedding_dim) """ self.eval() embeddings = [] with torch.no_grad(): for i in range(0, len(images), batch_size): batch_images = images[i:i + batch_size] # 预处理批次图片 batch_tensors = [] for img in batch_images: tensor = self.preprocess_image(img) batch_tensors.append(tensor.squeeze(0)) # 移除batch维度 # 堆叠成批次 batch_tensor = torch.stack(batch_tensors) # 移动到正确的设备 device = next(self.parameters()).device batch_tensor = batch_tensor.to(device) # 提取特征 batch_embeddings = self.forward(batch_tensor, normalize=normalize) embeddings.append(batch_embeddings.cpu().numpy()) return np.vstack(embeddings) @staticmethod def cosine_similarity(embedding1: np.ndarray, embedding2: np.ndarray) -> float: """ 计算两个embedding向量的余弦相似度 Args: embedding1: 第一个embedding向量 embedding2: 第二个embedding向量 Returns: float: 余弦相似度值 [-1, 1] """ # 确保向量是一维的 if len(embedding1.shape) > 1: embedding1 = embedding1.flatten() if len(embedding2.shape) > 1: embedding2 = embedding2.flatten() # 计算余弦相似度 dot_product = np.dot(embedding1, embedding2) norm1 = np.linalg.norm(embedding1) norm2 = np.linalg.norm(embedding2) if norm1 == 0 or norm2 == 0: return 0.0 return dot_product / (norm1 * norm2) @staticmethod def batch_cosine_similarity(query_embedding: np.ndarray, database_embeddings: np.ndarray) -> np.ndarray: """ 计算查询向量与数据库中所有向量的余弦相似度 Args: query_embedding: 查询向量,形状为 (embedding_dim,) database_embeddings: 数据库向量,形状为 (N, embedding_dim) Returns: np.ndarray: 相似度数组,形状为 (N,) """ # 确保查询向量是一维的 if len(query_embedding.shape) > 1: query_embedding = query_embedding.flatten() # 归一化向量(如果还没有归一化) query_norm = np.linalg.norm(query_embedding) if query_norm > 0: query_embedding = query_embedding / query_norm database_norms = np.linalg.norm(database_embeddings, axis=1) database_norms[database_norms == 0] = 1 # 避免除零 database_embeddings_normalized = database_embeddings / database_norms[:, np.newaxis] # 计算余弦相似度 similarities = np.dot(database_embeddings_normalized, query_embedding) return similarities def get_embedding_dim(self) -> int: """ 获取embedding向量维度 Returns: int: embedding维度 """ return self.embedding_dim def freeze_backbone(self): """ 冻结骨干网络参数,只训练embedding层 """ for param in self.backbone.parameters(): param.requires_grad = False def unfreeze_backbone(self): """ 解冻骨干网络参数,允许端到端训练 """ for param in self.backbone.parameters(): param.requires_grad = True def get_trainable_parameters(self): """ 获取可训练参数 Returns: generator: 可训练参数生成器 """ return filter(lambda p: p.requires_grad, self.parameters()) def create_resnet50_embedding(embedding_dim: int = 512, pretrained: bool = True, use_internal_preprocess: bool = False) -> ResNet50EmbeddingNet: """ 创建ResNet50 Embedding模型 Args: embedding_dim: embedding向量维度 pretrained: 是否使用预训练权重 use_internal_preprocess: 是否在forward中进行预处理 Returns: ResNet50EmbeddingNet: 网络模型实例 """ return ResNet50EmbeddingNet( embedding_dim=embedding_dim, pretrained=pretrained, use_internal_preprocess=use_internal_preprocess ) def create_mobile_resnet50_embedding(embedding_dim: int = 512) -> ResNet50EmbeddingNet: """ 创建移动端ResNet50 Embedding模型 Args: embedding_dim: embedding向量维度 Returns: ResNet50EmbeddingNet: 配置为移动端使用的网络模型实例 """ class MobileResNet50Embedding(ResNet50EmbeddingNet): """ 移动端ResNet50 Embedding模型 重写forward方法以包含预处理 """ def __init__(self, embedding_dim: int = 512): super().__init__(embedding_dim=embedding_dim, pretrained=True, use_internal_preprocess=False) def forward(self, x: torch.Tensor, normalize: bool = True) -> torch.Tensor: """ 移动端前向传播(自动包含预处理) Args: x: 输入tensor,形状为 [batch, 3, height, width],值范围 0-1 normalize: 是否对输出进行L2归一化 Returns: torch.Tensor: embedding向量 """ return self.forward_mobile(x, normalize=normalize) return MobileResNet50Embedding(embedding_dim=embedding_dim) if __name__ == "__main__": # 测试网络 print("=== ResNet50 Embedding网络测试 ===") # 创建模型 model = create_resnet50_embedding(embedding_dim=512, pretrained=True) print(f"模型参数数量: {sum(p.numel() for p in model.parameters() if p.requires_grad):,}") print(f"Embedding维度: {model.get_embedding_dim()}") # 测试前向传播 dummy_input = torch.randn(2, 3, 224, 224) embeddings = model(dummy_input) print(f"输入形状: {dummy_input.shape}") print(f"输出embedding形状: {embeddings.shape}") print(f"输出数值范围: [{embeddings.min():.3f}, {embeddings.max():.3f}]") # 测试L2归一化 embedding_norms = torch.norm(embeddings, p=2, dim=1) print(f"L2归一化后的向量模长: {embedding_norms}") # 测试图片预处理 try: import numpy as np # 创建一个测试图片 (RGB格式) test_image = Image.fromarray(np.random.randint(0, 255, (224, 224, 3), dtype=np.uint8)) # 测试单张图片特征提取 embedding = model.extract_embedding(test_image) print(f"单张图片embedding形状: {embedding.shape}") print(f"单张图片embedding模长: {np.linalg.norm(embedding):.3f}") # 测试批量特征提取 test_images = [test_image] * 3 batch_embeddings = model.extract_batch_embeddings(test_images, batch_size=2) print(f"批量embedding形状: {batch_embeddings.shape}") # 测试相似度计算 similarity = ResNet50EmbeddingNet.cosine_similarity(embedding, batch_embeddings[0]) print(f"自相似度: {similarity:.3f}") # 测试批量相似度计算 similarities = ResNet50EmbeddingNet.batch_cosine_similarity(embedding, batch_embeddings) print(f"批量相似度: {similarities}") except Exception as e: print(f"图片处理测试失败: {e}") # # 测试移动端模型 # print("\n=== 移动端模型测试 ===") # mobile_model = create_mobile_resnet50_embedding(embedding_dim=512) # mobile_input = torch.rand(1, 3, 256, 256) # 模拟移动端输入 # mobile_embedding = mobile_model(mobile_input) # print(f"移动端输入形状: {mobile_input.shape}") # print(f"移动端输出embedding形状: {mobile_embedding.shape}") # print(f"移动端embedding模长: {torch.norm(mobile_embedding, p=2, dim=1)}") # # print("\n=== 测试完成 ===")