Files
FoodClassifier/net/resnet_embedding.py
T

465 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)
)
# 图片预处理变换(用于推理)
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=== 测试完成 ===")