增加resnet_embedding的代码

This commit is contained in:
zhanghuan
2025-09-17 14:44:34 +08:00
parent 44031b92cd
commit 8818643025
+448
View File
@@ -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=== 测试完成 ===")