isLinXu
Pack FlorenceForge source for embedded HF Spaces deployment
e40db0e
Raw
History Blame Contribute Delete
17.4 kB
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
模型优化器
提供模型量化、剪枝、蒸馏等优化技术
"""
import torch
import torch.nn as nn
from typing import Optional, List, Dict, Any
import logging
import copy
import numpy as np
logger = logging.getLogger(__name__)
class DeploymentOptimizer:
"""部署模型优化器
提供多种模型优化技术(量化、剪枝、蒸馏等),专注于部署场景。
与 utils.optimization.ModelOptimizer(通用运行时优化工具)不同。
"""
def __init__(self, model: nn.Module):
"""初始化优化器
Args:
model: 要优化的模型
"""
self.original_model = model
self.optimized_model = None
def quantize_dynamic(
self,
qconfig_spec: Optional[Dict] = None,
dtype: torch.dtype = torch.qint8
) -> nn.Module:
"""动态量化
Args:
qconfig_spec: 量化配置
dtype: 量化数据类型
Returns:
量化后的模型
"""
try:
model_copy = copy.deepcopy(self.original_model)
model_copy.eval()
# 默认量化配置
if qconfig_spec is None:
qconfig_spec = {
nn.Linear: torch.quantization.default_dynamic_qconfig,
nn.LSTM: torch.quantization.default_dynamic_qconfig,
nn.GRU: torch.quantization.default_dynamic_qconfig
}
quantized_model = torch.quantization.quantize_dynamic(
model_copy,
qconfig_spec,
dtype=dtype
)
self.optimized_model = quantized_model
logger.info("动态量化完成")
return quantized_model
except Exception as e:
logger.error(f"动态量化失败: {e}")
raise
def quantize_static(
self,
calibration_data_loader,
qconfig: Optional[torch.quantization.QConfig] = None
) -> nn.Module:
"""静态量化
Args:
calibration_data_loader: 校准数据加载器
qconfig: 量化配置
Returns:
量化后的模型
"""
try:
model_copy = copy.deepcopy(self.original_model)
model_copy.eval()
# 设置量化配置
if qconfig is None:
qconfig = torch.quantization.get_default_qconfig('fbgemm')
model_copy.qconfig = qconfig
# 准备量化
torch.quantization.prepare(model_copy, inplace=True)
# 校准
logger.info("开始校准...")
with torch.no_grad():
for batch_idx, (data, _) in enumerate(calibration_data_loader):
model_copy(data)
if batch_idx >= 100: # 限制校准样本数量
break
# 转换为量化模型
quantized_model = torch.quantization.convert(model_copy, inplace=False)
self.optimized_model = quantized_model
logger.info("静态量化完成")
return quantized_model
except Exception as e:
logger.error(f"静态量化失败: {e}")
raise
def prune_unstructured(
self,
pruning_ratio: float = 0.2,
pruning_method: str = "magnitude"
) -> nn.Module:
"""非结构化剪枝
Args:
pruning_ratio: 剪枝比例
pruning_method: 剪枝方法
Returns:
剪枝后的模型
"""
try:
import torch.nn.utils.prune as prune
model_copy = copy.deepcopy(self.original_model)
# 收集要剪枝的参数
parameters_to_prune = []
for module in model_copy.modules():
if isinstance(module, (nn.Linear, nn.Conv2d)):
parameters_to_prune.append((module, 'weight'))
# 应用剪枝
if pruning_method == "magnitude":
prune.global_unstructured(
parameters_to_prune,
pruning_method=prune.L1Unstructured,
amount=pruning_ratio
)
elif pruning_method == "random":
prune.global_unstructured(
parameters_to_prune,
pruning_method=prune.RandomUnstructured,
amount=pruning_ratio
)
else:
raise ValueError(f"不支持的剪枝方法: {pruning_method}")
# 移除剪枝重参数化
for module, param_name in parameters_to_prune:
prune.remove(module, param_name)
self.optimized_model = model_copy
logger.info(f"非结构化剪枝完成,剪枝比例: {pruning_ratio}")
return model_copy
except ImportError:
logger.error("PyTorch版本不支持剪枝功能")
raise
except Exception as e:
logger.error(f"非结构化剪枝失败: {e}")
raise
def prune_structured(
self,
pruning_ratio: float = 0.2,
dim: int = 0
) -> nn.Module:
"""结构化剪枝
Args:
pruning_ratio: 剪枝比例
dim: 剪枝维度
Returns:
剪枝后的模型
"""
try:
import torch.nn.utils.prune as prune
model_copy = copy.deepcopy(self.original_model)
# 对每个线性层和卷积层进行结构化剪枝
for module in model_copy.modules():
if isinstance(module, (nn.Linear, nn.Conv2d)):
prune.ln_structured(
module,
name='weight',
amount=pruning_ratio,
n=2,
dim=dim
)
prune.remove(module, 'weight')
self.optimized_model = model_copy
logger.info(f"结构化剪枝完成,剪枝比例: {pruning_ratio}")
return model_copy
except ImportError:
logger.error("PyTorch版本不支持剪枝功能")
raise
except Exception as e:
logger.error(f"结构化剪枝失败: {e}")
raise
def knowledge_distillation(
self,
student_model: nn.Module,
teacher_model: nn.Module,
train_loader,
num_epochs: int = 10,
temperature: float = 4.0,
alpha: float = 0.7,
device: str = "cpu"
) -> nn.Module:
"""知识蒸馏
Args:
student_model: 学生模型
teacher_model: 教师模型
train_loader: 训练数据加载器
num_epochs: 训练轮数
temperature: 蒸馏温度
alpha: 蒸馏损失权重
device: 设备
Returns:
蒸馏后的学生模型
"""
try:
device = torch.device(device)
student_model = student_model.to(device)
teacher_model = teacher_model.to(device)
teacher_model.eval()
optimizer = torch.optim.Adam(student_model.parameters(), lr=1e-4)
criterion_ce = nn.CrossEntropyLoss()
criterion_kl = nn.KLDivLoss(reduction='batchmean')
logger.info(f"开始知识蒸馏训练,共{num_epochs}轮")
for epoch in range(num_epochs):
student_model.train()
total_loss = 0.0
for batch_idx, (data, target) in enumerate(train_loader):
data, target = data.to(device), target.to(device)
optimizer.zero_grad()
# 学生模型输出
student_output = student_model(data)
# 教师模型输出
with torch.no_grad():
teacher_output = teacher_model(data)
# 计算损失
# 硬标签损失
loss_ce = criterion_ce(student_output, target)
# 软标签损失(知识蒸馏)
loss_kl = criterion_kl(
torch.log_softmax(student_output / temperature, dim=1),
torch.softmax(teacher_output / temperature, dim=1)
) * (temperature ** 2)
# 总损失
loss = alpha * loss_kl + (1 - alpha) * loss_ce
loss.backward()
optimizer.step()
total_loss += loss.item()
if batch_idx % 100 == 0:
logger.info(
f"Epoch {epoch+1}/{num_epochs}, "
f"Batch {batch_idx}, Loss: {loss.item():.4f}"
)
avg_loss = total_loss / len(train_loader)
logger.info(f"Epoch {epoch+1} 平均损失: {avg_loss:.4f}")
self.optimized_model = student_model
logger.info("知识蒸馏完成")
return student_model
except Exception as e:
logger.error(f"知识蒸馏失败: {e}")
raise
def optimize_for_mobile(self) -> nn.Module:
"""移动端优化
Returns:
优化后的模型
"""
try:
from torch.utils.mobile_optimizer import optimize_for_mobile # noqa: F401
# 先转换为TorchScript
model_copy = copy.deepcopy(self.original_model)
model_copy.eval()
# 这里需要示例输入来trace模型
# 实际使用时需要提供合适的输入
logger.warning("移动端优化需要示例输入来trace模型")
# traced_model = torch.jit.trace(model_copy, example_input)
# optimized_model = optimize_for_mobile(traced_model)
# 暂时返回原模型
self.optimized_model = model_copy
logger.info("移动端优化完成(需要提供示例输入)")
return model_copy
except ImportError:
logger.error("移动端优化功能不可用")
raise
except Exception as e:
logger.error(f"移动端优化失败: {e}")
raise
def fuse_modules(self, modules_to_fuse: Optional[List[List[str]]] = None) -> nn.Module:
"""模块融合优化
Args:
modules_to_fuse: 要融合的模块列表
Returns:
融合后的模型
"""
try:
model_copy = copy.deepcopy(self.original_model)
model_copy.eval()
if modules_to_fuse is None:
# 自动检测可融合的模块
modules_to_fuse = self._detect_fusable_modules(model_copy)
if modules_to_fuse:
torch.quantization.fuse_modules(model_copy, modules_to_fuse, inplace=True)
logger.info(f"模块融合完成,融合了 {len(modules_to_fuse)} 组模块")
else:
logger.info("未发现可融合的模块")
self.optimized_model = model_copy
return model_copy
except Exception as e:
logger.error(f"模块融合失败: {e}")
raise
def _detect_fusable_modules(self, model: nn.Module) -> List[List[str]]:
"""自动检测可融合的模块
Args:
model: 模型
Returns:
可融合模块列表
"""
fusable_modules = []
# 简单的启发式检测
# 实际应用中可能需要更复杂的逻辑
for name, module in model.named_modules():
if isinstance(module, nn.Sequential):
submodules = list(module.children())
for i in range(len(submodules) - 1):
if (isinstance(submodules[i], nn.Conv2d) and
isinstance(submodules[i+1], nn.BatchNorm2d)):
fusable_modules.append([f"{name}.{i}", f"{name}.{i+1}"])
return fusable_modules
def compare_models(
self,
test_loader,
metrics: List[str] = None
) -> Dict[str, Dict[str, float]]:
"""比较原模型和优化模型的性能
Args:
test_loader: 测试数据加载器
metrics: 要计算的指标
Returns:
性能比较结果
"""
if self.optimized_model is None:
raise ValueError("尚未进行模型优化")
if metrics is None:
metrics = ['accuracy', 'inference_time', 'model_size']
results = {
'original': {},
'optimized': {},
'improvement': {}
}
# 评估原模型
original_metrics = self._evaluate_model(self.original_model, test_loader, metrics)
results['original'] = original_metrics
# 评估优化模型
optimized_metrics = self._evaluate_model(self.optimized_model, test_loader, metrics)
results['optimized'] = optimized_metrics
# 计算改进
for metric in metrics:
if metric in original_metrics and metric in optimized_metrics:
if metric == 'inference_time':
# 推理时间越小越好
improvement = (original_metrics[metric] - optimized_metrics[metric]) / original_metrics[metric] * 100
else:
# 其他指标越大越好
improvement = (optimized_metrics[metric] - original_metrics[metric]) / original_metrics[metric] * 100
results['improvement'][metric] = improvement
return results
def _evaluate_model(
self,
model: nn.Module,
test_loader,
metrics: List[str]
) -> Dict[str, float]:
"""评估模型性能
Args:
model: 要评估的模型
test_loader: 测试数据加载器
metrics: 要计算的指标
Returns:
性能指标
"""
import time
results = {}
model.eval()
if 'accuracy' in metrics:
correct = 0
total = 0
with torch.no_grad():
for data, target in test_loader:
output = model(data)
pred = output.argmax(dim=1, keepdim=True)
correct += pred.eq(target.view_as(pred)).sum().item()
total += target.size(0)
results['accuracy'] = correct / total
if 'inference_time' in metrics:
# 测量推理时间
times = []
with torch.no_grad():
for i, (data, _) in enumerate(test_loader):
if i >= 100: # 只测试前100个batch
break
start_time = time.time()
_ = model(data)
end_time = time.time()
times.append(end_time - start_time)
results['inference_time'] = np.mean(times)
if 'model_size' in metrics:
# 计算模型大小(参数数量)
total_params = sum(p.numel() for p in model.parameters())
results['model_size'] = total_params
return results
def get_optimization_summary(self) -> Dict[str, Any]:
"""获取优化摘要
Returns:
优化摘要信息
"""
if self.optimized_model is None:
return {"status": "未进行优化"}
original_params = sum(p.numel() for p in self.original_model.parameters())
optimized_params = sum(p.numel() for p in self.optimized_model.parameters())
summary = {
"original_parameters": original_params,
"optimized_parameters": optimized_params,
"parameter_reduction": (original_params - optimized_params) / original_params * 100,
"compression_ratio": original_params / optimized_params if optimized_params > 0 else float('inf')
}
return summary