From 881864302517ed07288e90de8b28dfa0be8c4916 Mon Sep 17 00:00:00 2001 From: zhanghuan <1262329256@qq.com> Date: Wed, 17 Sep 2025 14:44:34 +0800 Subject: [PATCH] =?UTF-8?q?=E5=A2=9E=E5=8A=A0resnet=5Fembedding=E7=9A=84?= =?UTF-8?q?=E4=BB=A3=E7=A0=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- net/resnet_embedding.py | 448 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 448 insertions(+) diff --git a/net/resnet_embedding.py b/net/resnet_embedding.py index e69de29..9f8175d 100644 --- a/net/resnet_embedding.py +++ b/net/resnet_embedding.py @@ -0,0 +1,448 @@ +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=== 测试完成 ===") \ No newline at end of file