Agent模型蒸馏与小模型部署策略
Agent模型蒸馏与小模型部署策略
引言
随着大语言模型在Agent系统中的广泛应用,一个日益突出的问题摆在了开发者面前:如何在资源受限的环境中部署具备Agent能力的模型。完整的大模型动辄需要数十GB的显存,推理延迟高,运营成本昂贵。而许多实际应用场景——如移动端助手、边缘设备智能体、低成本API服务——无法承受这样的资源消耗。
模型蒸馏技术为这一问题提供了可行的解决方案。通过将大模型(教师模型)的知识迁移到小模型(学生模型),可以在大幅降低模型参数量和资源需求的同时,尽可能保留教师模型的Agent能力。这不仅仅是简单的参数压缩,而是一套系统性的知识迁移方法论,涉及蒸馏数据构建、训练策略设计、能力评估和部署优化等多个环节。
本文将全面介绍Agent模型蒸馏的技术体系,从知识蒸馏的基本原理到小模型Agent能力的评估方法,再到边缘部署的工程实践,并给出完整的代码实现。
知识蒸馏方法
知识蒸馏的核心思想最早由Hinton等人在2015年系统化提出。其基本框架是:教师模型(大模型)生成软标签(soft labels)或中间表征,学生模型(小模型)在学习真实标签的同时,也学习教师模型的这些输出,从而获得教师模型蕴含的"暗知识"。
在Agent场景中,知识蒸馏比传统分类任务复杂得多。Agent能力不仅包括文本生成,还涉及工具调用、多步推理、上下文理解等复杂行为。因此,Agent蒸馏需要设计更丰富的蒸馏信号和更精细的训练策略。
Agent蒸馏的主要方法可以分为以下几类:基于响应的蒸馏(学习教师模型的输出分布)、基于特征的蒸馏(学习教师模型的中间层表征)、基于行为的蒸馏(学习教师模型的Agent行为模式)、以及基于轨迹的蒸馏(学习教师模型的多步推理轨迹)。
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Dict, List, Tuple
import copy
class AgentDistillationConfig:
"""Agent蒸馏配置"""
def __init__(
self,
temperature: float = 2.0,
alpha_logit: float = 0.5,
alpha_hidden: float = 0.3,
alpha_behavior: float = 0.2,
max_seq_len: int = 2048,
distill_layers: Optional[List[int]] = None
):
self.temperature = temperature
self.alpha_logit = alpha_logit
self.alpha_hidden = alpha_hidden
self.alpha_behavior = alpha_behavior
self.max_seq_len = max_seq_len
self.distill_layers = distill_layers or [0, 4, 8, 12]
class LogitDistillationLoss(nn.Module):
"""基于logit的蒸馏损失:学生模型学习教师模型的输出分布"""
def __init__(self, temperature: float = 2.0):
super().__init__()
self.temperature = temperature
def forward(
self,
student_logits: torch.Tensor,
teacher_logits: torch.Tensor,
labels: Optional[torch.Tensor] = None
) -> torch.Tensor:
# 软目标损失:KL散度
soft_loss = F.kl_div(
F.log_softmax(student_logits / self.temperature, dim=-1),
F.softmax(teacher_logits / self.temperature, dim=-1),
reduction="batchmean"
) * (self.temperature ** 2)
# 硬目标损失:交叉熵
if labels is not None:
hard_loss = F.cross_entropy(
student_logits.view(-1, student_logits.size(-1)),
labels.view(-1),
ignore_index=-100
)
return soft_loss + hard_loss
return soft_loss
class HiddenStateDistillationLoss(nn.Module):
"""基于隐藏状态的蒸馏损失:对齐中间层表征"""
def __init__(self, student_dim: int, teacher_dim: int):
super().__init__()
# 投影矩阵将学生维度映射到教师维度
self.projection = nn.Linear(student_dim, teacher_dim)
def forward(
self,
student_hidden: torch.Tensor,
teacher_hidden: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None
) -> torch.Tensor:
projected = self.projection(student_hidden)
if attention_mask is not None:
mask = attention_mask.unsqueeze(-1).expand_as(projected)
projected = projected * mask
teacher_hidden = teacher_hidden * mask
return F.mse_loss(projected, teacher_hidden)
class BehaviorDistillationLoss(nn.Module):
"""基于行为的蒸馏损失:对齐Agent行为模式"""
def __init__(self, num_action_types: int = 5):
super().__init__()
self.num_action_types = num_action_types
def forward(
self,
student_action_logits: torch.Tensor,
teacher_action_probs: torch.Tensor,
action_labels: torch.Tensor
) -> torch.Tensor:
# 行为类型分类损失
behavior_kl = F.kl_div(
F.log_softmax(student_action_logits, dim=-1),
teacher_action_probs,
reduction="batchmean"
)
behavior_ce = F.cross_entropy(student_action_logits, action_labels)
return behavior_kl + behavior_ce
class AgentDistillationTrainer:
"""Agent蒸馏训练器"""
def __init__(
self,
teacher_model: nn.Module,
student_model: nn.Module,
config: AgentDistillationConfig,
device: str = "cuda"
):
self.teacher = teacher_model
self.student = student_model
self.config = config
self.device = device
self.teacher.eval()
for p in self.teacher.parameters():
p.requires_grad = False
self.logit_loss = LogitDistillationLoss(config.temperature).to(device)
self.hidden_loss = None
self.behavior_loss = BehaviorDistillationLoss().to(device)
self.optimizer = torch.optim.AdamW(
student_model.parameters(), lr=1e-4, weight_decay=0.01
)
def train_step(self, batch: Dict[str, torch.Tensor]) -> Dict[str, float]:
"""执行一步蒸馏训练"""
input_ids = batch["input_ids"].to(self.device)
attention_mask = batch["attention_mask"].to(self.device)
labels = batch.get("labels")
if labels is not None:
labels = labels.to(self.device)
# 教师模型前向传播
with torch.no_grad():
teacher_outputs = self.teacher(
input_ids, attention_mask=attention_mask, output_hidden_states=True
)
teacher_logits = teacher_outputs.logits
teacher_hidden = teacher_outputs.hidden_states
# 学生模型前向传播
student_outputs = self.student(
input_ids, attention_mask=attention_mask, output_hidden_states=True
)
student_logits = student_outputs.logits
student_hidden = student_outputs.hidden_states
# 计算各部分损失
logit_loss = self.logit_loss(student_logits, teacher_logits, labels)
hidden_loss = torch.tensor(0.0, device=self.device)
if teacher_hidden and student_hidden:
for layer_idx in self.config.distill_layers:
if layer_idx < len(student_hidden) and layer_idx < len(teacher_hidden):
s_dim = student_hidden[layer_idx].size(-1)
t_dim = teacher_hidden[layer_idx].size(-1)
if self.hidden_loss is None or self.hidden_loss.projection.in_features != s_dim:
self.hidden_loss = HiddenStateDistillationLoss(s_dim, t_dim).to(self.device)
hidden_loss += self.hidden_loss(
student_hidden[layer_idx], teacher_hidden[layer_idx], attention_mask
)
total_loss = (
self.config.alpha_logit * logit_loss +
self.config.alpha_hidden * hidden_loss
)
# 反向传播
self.optimizer.zero_grad()
total_loss.backward()
torch.nn.utils.clip_grad_norm_(self.student.parameters(), 1.0)
self.optimizer.step()
return {
"total_loss": total_loss.item(),
"logit_loss": logit_loss.item(),
"hidden_loss": hidden_loss.item() if isinstance(hidden_loss, torch.Tensor) else hidden_loss,
}
蒸馏方法的选择取决于具体场景。对于以文本生成为主的Agent,logit蒸馏是最核心的方法;对于需要深度推理的Agent,隐藏状态蒸馏能更好地传递教师的推理能力;对于工具调用密集的Agent,行为蒸馏可以确保学生模型学会正确的工具使用模式。实际应用中,多种蒸馏方法组合使用效果最佳。
教师-学生架构设计
教师-学生架构的设计是蒸馏成功的关键因素之一。教师模型的选择、学生模型的架构、以及两者之间的容量比都会影响蒸馏效果。
教师模型应该选择在目标任务上表现最好的大模型。在Agent场景中,教师模型需要具备强大的工具调用能力、多步推理能力和指令遵循能力。学生模型的选择则需要综合考虑部署约束和蒸馏可行性。学生模型太小可能导致无法吸收教师的知识,太大则失去了蒸馏的意义。
一个重要的实践原则是:学生模型与教师模型最好属于同一模型族。同族模型的词表、tokenizer、架构风格相似,蒸馏时知识迁移的效率更高。例如,如果教师是7B参数的模型,学生可以选择1.5B或3B的同族模型。
import torch.nn as nn
from typing import Optional, List
class TeacherStudentConfig:
"""教师-学生架构配置"""
def __init__(self):
self.teacher_config = {
"model_name": "large-agent-7b",
"num_layers": 32,
"hidden_size": 4096,
"num_heads": 32,
"vocab_size": 32000,
"max_seq_len": 4096
}
self.student_config = {
"model_name": "small-agent-1.5b",
"num_layers": 12,
"hidden_size": 2048,
"num_heads": 16,
"vocab_size": 32000,
"max_seq_len": 2048
}
# 层映射:学生第i层对应教师第j层
self.layer_mapping = self._build_layer_mapping()
def _build_layer_mapping(self) -> List[tuple]:
"""构建学生层到教师层的映射"""
t_layers = self.teacher_config["num_layers"]
s_layers = self.student_config["num_layers"]
mapping = []
for s_idx in range(s_layers):
t_idx = int(s_idx * (t_layers - 1) / max(s_layers - 1, 1))
mapping.append((s_idx, t_idx))
return mapping
def get_capacity_ratio(self) -> float:
"""计算学生/教师容量比"""
t_params = (self.teacher_config["num_layers"] * self.teacher_config["hidden_size"] ** 2)
s_params = (self.student_config["num_layers"] * self.student_config["hidden_size"] ** 2)
return s_params / t_params
class AdaptiveProjection(nn.Module):
"""自适应投影层:处理教师和学生之间的维度差异"""
def __init__(self, student_dim: int, teacher_dim: int, num_heads: int = 8):
super().__init__()
self.student_dim = student_dim
self.teacher_dim = teacher_dim
self.num_heads = num_heads
self.head_dim = student_dim // num_heads
self.linear = nn.Linear(student_dim, teacher_dim)
self.layer_norm = nn.LayerNorm(teacher_dim)
self.gate = nn.Sequential(
nn.Linear(teacher_dim * 2, teacher_dim),
nn.ReLU(),
nn.Linear(teacher_dim, 1),
nn.Sigmoid()
)
def forward(self, student_hidden: torch.Tensor, teacher_hidden: torch.Tensor) -> torch.Tensor:
"""将学生隐藏状态投影到教师空间,并计算门控权重"""
projected = self.linear(student_hidden)
projected = self.layer_norm(projected)
# 门控机制:决定哪些位置需要更紧密地对齐
gate_input = torch.cat([projected, teacher_hidden], dim=-1)
gate_weight = self.gate(gate_input)
# 加权对齐
aligned = gate_weight * projected + (1 - gate_weight) * teacher_hidden.detach()
return aligned, gate_weight
class ProgressiveDistillationScheduler:
"""渐进式蒸馏调度器:分阶段逐步蒸馏"""
def __init__(self, trainer: AgentDistillationTrainer, stages: List[dict]):
self.trainer = trainer
self.stages = stages
self.current_stage = 0
def get_current_config(self) -> dict:
"""获取当前阶段的配置"""
return self.stages[self.current_stage]
def should_advance(self, metrics: dict) -> bool:
"""判断是否应该进入下一阶段"""
config = self.get_current_config()
if metrics.get("eval_score", 0) >= config.get("target_score", 0.9):
return True
if metrics.get("epoch", 0) >= config.get("max_epochs", 10):
return True
return False
def advance_stage(self):
"""进入下一阶段"""
if self.current_stage < len(self.stages) - 1:
self.current_stage += 1
config = self.get_current_config()
# 调整训练参数
for param_group in self.trainer.optimizer.param_groups:
param_group["lr"] = config.get("lr", 1e-4)
# 调整蒸馏权重
self.trainer.config.alpha_logit = config.get("alpha_logit", 0.5)
self.trainer.config.alpha_hidden = config.get("alpha_hidden", 0.3)
return True
return False
渐进式蒸馏是一种有效的训练策略。它将蒸馏过程分为多个阶段,早期阶段侧重于学习教师模型的基本语言能力,后期阶段逐步引入更复杂的Agent行为蒸馏。这种分阶段的方法可以避免学生模型在训练初期被过于复杂的蒸馏信号淹没,从而获得更稳定的训练过程。
蒸馏数据构建
蒸馏数据的质量直接决定了学生模型的最终能力。在Agent场景中,蒸馏数据的构建比传统NLP任务复杂得多,需要覆盖多种Agent行为模式。
高质量的蒸馏数据应该包含:多轮对话数据(覆盖上下文理解能力)、工具调用数据(覆盖工具选择和参数生成能力)、推理链数据(覆盖思维链推理能力)、错误纠正数据(覆盖自我反思和修正能力)、以及多任务混合数据(覆盖任务泛化能力)。
数据的构建可以通过教师模型自动生成,也可以利用真实交互日志。自动生成的好处是可以大规模生产,但需要设计多样化的提示模板来确保数据覆盖面。真实交互日志更贴近实际使用场景,但数量有限且可能包含噪声。
import json
import random
from typing import List, Dict, Optional
from dataclasses import dataclass
@dataclass
class AgentDistillationSample:
"""Agent蒸馏样本"""
conversation: List[Dict[str, str]]
teacher_response: str
teacher_logits: Optional[List[float]] = None
teacher_hidden_summary: Optional[Dict] = None
action_type: str = "text" # text, tool_call, reasoning, reflection
tool_name: Optional[str] = None
metadata: Dict = field(default_factory=dict)
class DistillationDataBuilder:
"""蒸馏数据构建器"""
def __init__(self, teacher_model, tokenizer, max_samples: int = 10000):
self.teacher = teacher_model
self.tokenizer = tokenizer
self.max_samples = max_samples
self.prompt_templates = self._init_prompt_templates()
def _init_prompt_templates(self) -> Dict[str, List[str]]:
"""初始化多样化的提示模板"""
return {
"tool_call": [
"你是一个智能助手,请使用合适的工具完成以下任务:{task}",
"作为Agent,你需要调用工具来解决这个问题:{task}",
"请分析以下需求并选择适当的工具:{task}",
],
"multi_turn": [
"用户:{turn1}\n助手:{response1}\n用户:{turn2}\n请继续对话。",
"对话历史:{history}\n用户最新消息:{message}\n请回复。",
],
"reasoning": [
"请逐步分析并解决以下问题:{problem}",
"使用思维链推理回答:{question}",
"请分解以下复杂任务并逐步执行:{task}",
],
"reflection": [
"上一次尝试失败了,原因是:{error}。请分析并给出改进方案。",
"你的回答有以下问题:{critique}。请修正。",
],
}
def generate_samples(self, task_pool: List[str]) -> List[AgentDistillationSample]:
"""生成蒸馏样本"""
samples = []
categories = list(self.prompt_templates.keys())
for task in task_pool[:self.max_samples]:
category = random.choice(categories)
template = random.choice(self.prompt_templates[category])
prompt = template.format(task=task, problem=task, question=task)
# 教师模型生成
teacher_response = self._generate_teacher_response(prompt)
sample = AgentDistillationSample(
conversation=[{"role": "user", "content": prompt}],
teacher_response=teacher_response,
action_type=category,
)
samples.append(sample)
return samples
def _generate_teacher_response(self, prompt: str) -> str:
"""使用教师模型生成响应"""
inputs = self.tokenizer(prompt, return_tensors="pt")
with torch.no_grad():
outputs = self.teacher.generate(
**inputs, max_new_tokens=512, do_sample=True,
temperature=0.7, top_p=0.9
)
return self.tokenizer.decode(outputs[0], skip_special_tokens=True)
def augment_with_tool_traces(self, base_samples: List[AgentDistillationSample],
tool_definitions: List[Dict]) -> List[AgentDistillationSample]:
"""用工具调用轨迹增强数据"""
augmented = []
for sample in base_samples:
for tool in random.sample(tool_definitions, min(3, len(tool_definitions))):
tool_prompt = f"可用工具:{json.dumps(tool, ensure_ascii=False)}\n任务:{sample.conversation[0]['content']}"
teacher_response = self._generate_teacher_response(tool_prompt)
new_sample = AgentDistillationSample(
conversation=[{"role": "user", "content": tool_prompt}],
teacher_response=teacher_response,
action_type="tool_call",
tool_name=tool.get("name"),
)
augmented.append(new_sample)
return base_samples + augmented
def build_quality_filter(self) -> callable:
"""构建数据质量过滤器"""
def filter_fn(sample: AgentDistillationSample) -> bool:
if len(sample.teacher_response) < 10:
return False
if len(sample.teacher_response) > 2000:
return False
if sample.teacher_response.count("{") > 20:
return False
return True
return filter_fn
def to_training_format(self, samples: List[AgentDistillationSample]) -> List[Dict]:
"""转换为训练格式"""
training_data = []
for sample in samples:
full_text = sample.conversation[0]["content"] + "\n" + sample.teacher_response
tokenized = self.tokenizer(full_text, truncation=True,
max_length=2048, return_tensors="pt")
training_data.append({
"input_ids": tokenized["input_ids"],
"attention_mask": tokenized["attention_mask"],
"labels": tokenized["input_ids"].clone(),
"action_type": sample.action_type,
})
return training_data
数据构建中的一个关键问题是数据多样性。如果所有蒸馏样本都来自同一类型的任务,学生模型会过拟合到这一类型,丧失泛化能力。通过设计多样化的提示模板、混合不同类型的任务、引入随机扰动等方式,可以有效提升数据多样性。此外,数据质量过滤也是必要的,过短、过长、格式异常的样本应该被过滤掉。
小模型Agent能力评估
蒸馏后的学生模型需要经过严格的评估才能投入部署。传统的NLP评估指标如BLEU、ROUGE等无法全面衡量Agent能力。Agent能力评估需要覆盖多个维度:指令遵循能力、工具调用准确性、多步推理正确性、上下文理解能力、错误恢复能力等。
评估方法可以分为自动评估和人工评估。自动评估使用预定义的测试集和评估脚本,可以快速得到量化指标。人工评估由人类评估者对模型输出进行打分,更贴近实际体验但成本较高。在实际项目中,两种方法结合使用效果最佳。
import json
from typing import List, Dict, Optional
from dataclasses import dataclass
import re
@dataclass
class AgentEvalResult:
"""Agent评估结果"""
task_id: str
instruction_following: float
tool_call_accuracy: float
reasoning_correctness: float
context_understanding: float
error_recovery: float
overall_score: float
details: Dict
class AgentCapabilityEvaluator:
"""Agent能力评估器"""
def __init__(self, test_cases: List[Dict]):
self.test_cases = test_cases
self.tool_pattern = re.compile(r'"tool"\s*:\s*"(\w+)"')
self.param_pattern = re.compile(r'"(\w+)"\s*:\s*"([^"]*)"')
def evaluate_model(self, model, tokenizer) -> List[AgentEvalResult]:
"""评估模型的Agent能力"""
results = []
for case in self.test_cases:
result = self._evaluate_single_case(model, tokenizer, case)
results.append(result)
return results
def _evaluate_single_case(self, model, tokenizer, case: Dict) -> AgentEvalResult:
"""评估单个测试用例"""
prompt = case["prompt"]
expected = case.get("expected_response", "")
expected_tool = case.get("expected_tool")
expected_params = case.get("expected_params", {})
# 模型生成
inputs = tokenizer(prompt, return_tensors="pt")
with torch.no_grad():
outputs = model.generate(**inputs, max_new_tokens=512, do_sample=False)
response = tokenizer.decode(outputs[0], skip_special_tokens=True)
# 各维度评估
instr_score = self._eval_instruction_following(response, case.get("constraints", []))
tool_score = self._eval_tool_call(response, expected_tool)
reasoning_score = self._eval_reasoning(response, case.get("reasoning_steps", []))
context_score = self._eval_context(response, case.get("context_keys", []))
recovery_score = self._eval_error_recovery(response, case.get("error_scenario"))
overall = (instr_score * 0.2 + tool_score * 0.25 + reasoning_score * 0.25 +
context_score * 0.15 + recovery_score * 0.15)
return AgentEvalResult(
task_id=case["id"],
instruction_following=instr_score,
tool_call_accuracy=tool_score,
reasoning_correctness=reasoning_score,
context_understanding=context_score,
error_recovery=recovery_score,
overall_score=overall,
details={"response": response[:500], "expected": expected[:500]}
)
def _eval_instruction_following(self, response: str, constraints: List[str]) -> float:
"""评估指令遵循能力"""
if not constraints:
return 1.0
satisfied = 0
for constraint in constraints:
if constraint.lower() in response.lower():
satisfied += 1
return satisfied / len(constraints)
def _eval_tool_call(self, response: str, expected_tool: Optional[str]) -> float:
"""评估工具调用准确性"""
if expected_tool is None:
return 1.0
tool_match = self.tool_pattern.search(response)
if tool_match and tool_match.group(1) == expected_tool:
return 1.0
return 0.0
def _eval_reasoning(self, response: str, reasoning_steps: List[str]) -> float:
"""评估推理正确性"""
if not reasoning_steps:
return 1.0
matched = 0
for step in reasoning_steps:
if any(keyword in response.lower() for keyword in step.lower().split()):
matched += 1
return matched / len(reasoning_steps)
def _eval_context(self, response: str, context_keys: List[str]) -> float:
"""评估上下文理解能力"""
if not context_keys:
return 1.0
found = sum(1 for key in context_keys if key.lower() in response.lower())
return found / len(context_keys)
def _eval_error_recovery(self, response: str, error_scenario: Optional[Dict]) -> float:
"""评估错误恢复能力"""
if error_scenario is None:
return 1.0
if "apology" in response.lower() or "抱歉" in response:
return 0.5
if "alternative" in response.lower() or "替代" in response:
return 1.0
return 0.0
def generate_report(self, results: List[AgentEvalResult]) -> str:
"""生成评估报告"""
n = len(results)
report = "Agent能力评估报告\n"
report += f"测试用例数: {n}\n"
report += f"指令遵循: {sum(r.instruction_following for r in results)/n:.2%}\n"
report += f"工具调用: {sum(r.tool_call_accuracy for r in results)/n:.2%}\n"
report += f"推理正确: {sum(r.reasoning_correctness for r in results)/n:.2%}\n"
report += f"上下文理解: {sum(r.context_understanding for r in results)/n:.2%}\n"
report += f"错误恢复: {sum(r.error_recovery for r in results)/n:.2%}\n"
report += f"综合得分: {sum(r.overall_score for r in results)/n:.2%}\n"
# 找出弱项
avg_scores = {
"指令遵循": sum(r.instruction_following for r in results)/n,
"工具调用": sum(r.tool_call_accuracy for r in results)/n,
"推理正确": sum(r.reasoning_correctness for r in results)/n,
"上下文理解": sum(r.context_understanding for r in results)/n,
"错误恢复": sum(r.error_recovery for r in results)/n,
}
weakest = min(avg_scores, key=avg_scores.get)
report += f"最弱项: {weakest} ({avg_scores[weakest]:.2%})\n"
return report
评估结果的分析和反馈是持续改进的关键。通过识别学生模型的弱项,可以针对性地补充蒸馏数据、调整蒸馏权重、增加特定类型的训练样本。这种迭代式的评估-改进循环是提升小模型Agent能力的有效方法。
边缘部署
小模型蒸馏完成后,边缘部署是最终落地环节。边缘部署面临的主要挑战包括:计算资源有限、内存受限、功耗约束、网络连接不稳定等。需要针对这些约束进行专门的部署优化。
部署优化主要包括:模型格式转换(如转换为ONNX、TensorRT等格式)、算子融合、动态量化、模型分片加载、以及针对特定硬件的优化。
import torch
import torch.nn as nn
from typing import Dict, Optional, Tuple
import json
import os
class EdgeDeploymentOptimizer:
"""边缘部署优化器"""
def __init__(self, model: nn.Module, tokenizer):
self.model = model
self.tokenizer = tokenizer
self.optimization_log = []
def apply_dynamic_quantization(self) -> nn.Module:
"""应用动态量化(适用于CPU部署)"""
quantized_model = torch.quantization.quantize_dynamic(
self.model,
{nn.Linear, nn.LayerNorm},
dtype=torch.qint8
)
original_size = sum(p.nelement() * p.element_size() for p in self.model.parameters())
quantized_size = sum(p.nelement() * p.element_size() for p in quantized_model.parameters())
self.optimization_log.append({
"optimization": "dynamic_quantization",
"original_size_mb": original_size / 1024 / 1024,
"optimized_size_mb": quantized_size / 1024 / 1024,
"compression_ratio": original_size / max(quantized_size, 1)
})
return quantized_model
def apply_pruning(self, sparsity: float = 0.3) -> nn.Module:
"""应用结构化剪枝"""
parameters_to_prune = []
for name, module in self.model.named_modules():
if isinstance(module, nn.Linear):
parameters_to_prune.append((module, "weight"))
torch.nn.utils.prune.global_unstructured(
parameters_to_prune,
pruning_method=torch.nn.utils.prune.L1Unstructured,
amount=sparsity
)
for module, _ in parameters_to_prune:
torch.nn.utils.prune.remove(module, "weight")
self.optimization_log.append({
"optimization": "pruning",
"sparsity": sparsity,
})
return self.model
def export_to_onnx(self, output_path: str, max_seq_len: int = 512):
"""导出为ONNX格式"""
self.model.eval()
dummy_input = torch.randint(0, 32000, (1, max_seq_len), dtype=torch.long)
torch.onnx.export(
self.model,
dummy_input,
output_path,
export_params=True,
opset_version=14,
do_constant_folding=True,
input_names=["input_ids"],
output_names=["logits"],
dynamic_axes={
"input_ids": {0: "batch", 1: "sequence"},
"logits": {0: "batch", 1: "sequence"}
}
)
self.optimization_log.append({
"optimization": "onnx_export",
"output_path": output_path,
"max_seq_len": max_seq_len
})
def benchmark_inference(self, test_prompts: List[str], num_warmup: int = 3) -> Dict:
"""基准测试推理性能"""
import time
latencies = []
for _ in range(num_warmup):
inputs = self.tokenizer(test_prompts[0], return_tensors="pt")
with torch.no_grad():
_ = self.model(**inputs)
for prompt in test_prompts:
inputs = self.tokenizer(prompt, return_tensors="pt")
start = time.perf_counter()
with torch.no_grad():
outputs = self.model.generate(**inputs, max_new_tokens=128)
latency = time.perf_counter() - start
latencies.append(latency)
return {
"avg_latency_ms": sum(latencies) / len(latencies) * 1000,
"p50_latency_ms": sorted(latencies)[len(latencies)//2] * 1000,
"p95_latency_ms": sorted(latencies)[int(len(latencies)*0.95)] * 1000,
"max_latency_ms": max(latencies) * 1000,
}
class ModelPartitioner:
"""模型分片加载器:用于内存受限设备"""
def __init__(self, model_path: str, num_partitions: int = 3):
self.model_path = model_path
self.num_partitions = num_partitions
self.partitions = []
def partition_model(self, model: nn.Module):
"""将模型按层分片"""
layers = list(model.named_children())
partition_size = len(layers) // self.num_partitions
for i in range(self.num_partitions):
start = i * partition_size
end = start + partition_size if i < self.num_partitions - 1 else len(layers)
partition = nn.Sequential(*[layer for _, layer in layers[start:end]])
partition_path = os.path.join(self.model_path, f"partition_{i}.pt")
torch.save(partition.state_dict(), partition_path)
self.partitions.append({
"path": partition_path,
"layers": [name for name, _ in layers[start:end]],
"size_mb": os.path.getsize(partition_path) / 1024 / 1024
})
def load_partition(self, index: int, model_class, config) -> nn.Module:
"""按需加载单个分片"""
partition_info = self.partitions[index]
partition = model_class(config)
partition.load_state_dict(torch.load(partition_info["path"]))
return partition
边缘部署还需要考虑推理框架的选择。ONNX Runtime适合跨平台部署,TensorRT适合NVIDIA GPU优化,CoreML适合Apple设备,TFLite适合Android设备。根据目标平台选择合适的推理框架,可以最大化小模型的推理效率。
在实际的Agent边缘部署中,还需要设计合理的缓存策略和降级方案。当设备资源不足时,可以降级为更简单的模型或回退到云端推理。当网络连接中断时,本地缓存的历史对话和常用工具定义可以保证基本的Agent功能。这些工程细节虽然不涉及模型本身,但对用户体验的影响至关重要。
最后,蒸馏小模型的持续迭代也是不可忽视的。随着教师模型的升级和新任务的出现,学生模型也需要定期重新蒸馏和优化。建立自动化的蒸馏-评估-部署流水线,可以大幅降低维护成本,确保小模型始终保持良好的Agent能力。
- 点赞
- 收藏
- 关注作者
评论(0)