增加resnet_embedding的代码
This commit is contained in:
@@ -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=== 测试完成 ===")
|
||||||
Reference in New Issue
Block a user