🚀 Hugging Face 生态接入规范
为了让非官方的 engram-peft 模型能无缝对接 HF 的主流工具,修改接口层时必须注意以下红线:
1. 适配 generate() 文本生成
Engram 是一种序列依赖的架构,在自回归生成时,缓存机制至关重要。
- 必须确保模型正确实现了
prepare_inputs_for_generation方法。 - 除了标准的
input_ids和past_key_values,如果 Engram 需要传递特定的 memory index 或 mask,必须在kwargs中正确透传。 - 注意
past_key_values格式(通常是DynamicCache或 Tuple),不要破坏其原生结构。 - 必须支持
use_cache=True和use_cache=False两种模式。
2. 适配 HF Trainer
Trainer 对模型的前向传播返回值有严格要求:
- 训练模式下,
forward()的返回值必须包含loss字段,且通常建议返回CausalLMOutputWithPast。 - 必须处理好
labels参数的偏移 (shift) 逻辑(即inputs和labels错位计算 CrossEntropy)。如果你的代码直接覆盖了原生的forward,切记补齐损失计算逻辑。 - 必须支持
gradient_checkpointing以节省显存。
3. 常见坑与注意事项
Trainer会自动将模型移动到正确的设备,不要在forward中手动调用.to(device)- 确保所有自定义张量都正确跟随模型的设备和 dtype
- 对于稀疏张量,必须处理好
Trainer的梯度累积逻辑
Source: QingGo/engram-peft — distributed by TomeVault.