466 lines
17 KiB
Python
466 lines
17 KiB
Python
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 (使用新的weights参数)
|
||
if pretrained:
|
||
# V2效果会好一点
|
||
self.backbone = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)
|
||
else:
|
||
self.backbone = models.resnet50(weights=None)
|
||
|
||
# 获取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)
|
||
# )
|
||
# 暂时不用加Dropout
|
||
self.embedding_layer = nn.Sequential(
|
||
nn.Linear(backbone_output_dim, embedding_dim),
|
||
nn.BatchNorm1d(embedding_dim),
|
||
nn.ReLU(inplace=True)
|
||
)
|
||
|
||
# 图片预处理变换(用于推理)
|
||
# TODO 这里应该要修改归一化的逻辑,不然不一样。
|
||
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]
|
||
"""
|
||
# # 通过 scale_factor 避免对 size 的整数检查(规避 RecursionError)
|
||
# h, w = x.shape[2], x.shape[3]
|
||
# # 防御:避免除零
|
||
# if h == 0 or w == 0:
|
||
# raise ValueError(f"Invalid input size: height={h}, width={w}")
|
||
# scale_h = 224.0 / float(h)
|
||
# scale_w = 224.0 / float(w)
|
||
#
|
||
# x = F.interpolate(x, scale_factor=(scale_h, scale_w),
|
||
# mode='bilinear', align_corners=False)
|
||
# 缩放到 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归一化(用于余弦相似度计算),p=2表示L2范数(即平方和再开根号),dim=1表示在第二个维度上做归一化
|
||
# 形状是[batch_size, embedding_dim]
|
||
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)
|
||
|
||
# 显式调用基类 forward,避免子类重写的 forward 形成递归
|
||
return ResNet50EmbeddingNet.forward(self, 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)
|
||
# embedding = self.forward(x, normalize=False)
|
||
|
||
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,pretrained: bool = True) -> 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:
|
||
"""
|
||
移动端前向传播(自动包含预处理)
|
||
"""
|
||
# 在子类里做预处理,然后显式调用基类 forward
|
||
x = self.mobile_preprocess(x)
|
||
return ResNet50EmbeddingNet.forward(self, 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=== 测试完成 ===") |