AReaL 实战指南:如何为 Archon 训练引擎接入新的 HuggingFace 模型架构

AReaL 实战指南:如何为 Archon 训练引擎接入新的 HuggingFace 模型架构 AReaL 实战指南如何为 Archon 训练引擎接入新的 HuggingFace 模型架构【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL本文为 AReaLThe RL Bridge for LLM-based Agent Applications中 Archon 训练引擎的模型扩展指南。读完你将掌握一条完整的落地路径从分析 HuggingFace 目标模型的架构特征Attention 变体、FFN、MoE、RoPE、归一化到基于 qwen2/qwen3 参考实现生成args.py、model.py、state_dict_adapter.py、parallelize.py等核心文件最终通过ModelSpec注册机制让新模型类型被引擎自动识别并按仓库既有测试分层完成从 CPU 单测到单卡精度对齐的验证。该指南的原始文档是仓库中的 Agent 技能文件 SKILL.md本文在其骨架上结合当前仓库源码逐一印证每个契约contract与实现细节。何时使用本流程与前置条件本流程适用于以下场景引自原文档 When to Use询问如何给 Archon 添加一个新模型希望让ArchonEngine支持一个新的模型家族例如 Llama、Mistral、DeepSeek 等 decoder-only 架构需要为 Archon 增加一个新的ModelSpec或模型类型。开始前需确认三个前置条件目标模型在 HuggingFace 上可用且其config.json中带有model_type字段——这个字符串正是引擎注册表查找的键你已知道目标模型的 HuggingFace 模型 ID例如meta-llama/Llama-3-8B模型使用标准 decoder-only Transformer 架构。从源码结构看model_type到ModelSpec的映射由一个进程内注册表维护注册与查询逻辑集中在 model_spec.pyregister_model_spec(spec)遍历spec.supported_model_types若某个model_type已被注册会直接抛出ValueError见 model_spec.py#L101-L110因此一个model_type全局只能归属一个模型实现get_model_spec(model_type)对未注册的类型抛出KeyError并列出当前所有可用类型is_supported_model(model_type)与get_supported_model_types()分别用于能力探测与枚举。总体流程与目标文件骨架整个接入流程分为 10 步分析目标模型架构 → 选择参考实现 → 实现args.py→ 实现model.py→ 实现rope.py→ 实现state_dict_adapter.py→ 实现parallelize.py→ 编写spec.py并注册 → 在包级__init__.py中导入 → 分层验证与测试。每个模型目录的标准骨架如下对应 Step 2 给出的复制模板areal/experimental/models/archon/model/ __init__.py spec.py model/ args.py model.py rope.py state_dict_adapter.py infra/ parallelize.py当前仓库中已存在三个完整参考实现均符合该骨架qwen2、qwen3 与 qwen3_5。后两者的存在也印证了原文档维护者注释中的约定新增参考模型时应同步更新技能文档的参考表——从源码结构看archon 包的init.py 中已有qwen3_5的 spec 导入行说明该机制可承载多模态复合命名空间等更复杂的扩展其 state dict adapter 通过覆写_maybe_composite_hf_key钩子把文本权重映射到多模态检查点命名空间见 base.py#L123-L130。Step 1分析目标模型架构第一步是阅读 HuggingFace 模型的源码抽取关键架构信息形成一份明确的特征清单。1读config.json等价于AutoConfig.from_pretrainedmodel_type字符串注册表查找的键全部架构超参数hidden_size、num_layers等模型特有字段如qk_norm、attention_bias、MoE 相关字段。2读 HuggingFace 的modeling_*.py确认以下维度Attention 变体是否有 Q/K norm是否有 attention bias是否滑动窗口是否多潜在注意力MLAFFN 变体SwiGLUgate_proj up_proj down_projGeGLU标准 MLPMoE 支持是否 MoE 层何种 router是否有 shared expertsRoPE 变体标准 RoPE / YaRN / NTK-awareinv_freq的公式是什么归一化RMSNorm 还是 LayerNormPre-norm 还是 post-norm是否 elementwise affine权重共享config 中是否出现tie_word_embeddingsState dict 键名HF 权重的命名约定是什么3产出检查单原文档给出的模板Target model: name HF model_type: model_type (and variants like model_type_moE if applicable) Attention: [standard GQA / with QK norm / with bias / sliding window / ...] FFN: [SwiGLU / GeGLU / standard MLP / ...] MoE: [no / yes - num_experts, top_k, shared_experts] RoPE: [standard / YaRN / NTK-aware / ...] Norm: [RMSNorm / LayerNorm] with [pre-norm / post-norm] Weight tying: [yes / no]这份清单直接决定后续每一步的分支选择尤其是 Step 2 的参考模型选择与 Step 7 的 TP plan 设计。Step 2选择参考模型根据目标模型特征选择最接近的现有实现作为起点目标特征参考实现选择理由纯 Dense、标准 GQA、无 QK normqwen2最简单基线纯 dense 结构有 QK norm或有 MoEqwen3支持 QK norm MoE shared experts随后复制参考模型目录作为新模型的起点骨架见上节。从源码结构看qwen3之所以能同时承载 dense 与 MoE 两种形态是因为其 Qwen3ModelArgs 内置了完整的 MoE 配置域moe_enabled、moe_inter_dim、moe_args复用 moe/args.py 中的MoEArgs以及decoder_sparse_stepMoE 层间隔1 表示每层都是 MoE2 表示隔层 MoE0/负数表示禁用。因此若你的目标模型带 MoE直接以 qwen3 为参考可省去大量 MoE 相关代码。Step 3实现args.py—— HF config 到 ModelArgs 的映射ModelModelArgs是连接 HuggingFace 配置与 Archon 模型实现的桥梁它必须继承BaseModelArgs定义于 base.py。基类契约dataclass class ModelModelArgs(BaseModelArgs): # ... 模型特有字段 ... classmethod def from_hf_config( cls, hf_config: PretrainedConfig, is_critic: bool False, **kwargs, ) - ModelModelArgs: # 将 HF config 字段映射为 Archon model args ...两个关键细节值得注意对照基类实现 base.py#L26-L54BaseModelArgs自带字段attn_type默认varlen可选sdpa对应 Archon 中两种注意力后端attention/sdpa.py 与 attention/varlen.py。子类在from_hf_config中通过kwargs.get(attn_type, BaseModelArgs.attn_type)传递这一点在 Qwen2ModelArgs 与 Qwen3ModelArgs 中都是这样实现的rope_theta的提取必须兼容 transformers 大版本差异v4 中它是 config 的直接属性v5 中移入了rope_parameters字典。基类已提供_get_rope_theta(hf_config, default)辅助方法base.py#L33-L45新模型应直接复用而不要自己重写。字段约定引自原文档字段名遵循 Archon 惯例dim、n_layers、n_heads、n_kv_heads、vocab_size、head_dim、hidden_dim、norm_eps、rope_theta等默认值应对应目标模型的最小规格为模型特有特性追加字段如attention_bias、qk_norm、sliding_windowfrom_hf_config()中对可选字段一律使用getattr(hf_config, field_name, default)并处理变体专属字段如仅 MoE 变体才有的字段。以 Qwen2ModelArgs.from_hf_config 为例可以看到典型的映射写法return cls( dimhf_config.hidden_size, n_layershf_config.num_hidden_layers, n_headshf_config.num_attention_heads, n_kv_headsgetattr( hf_config, num_key_value_heads, hf_config.num_attention_heads ), # 无 GQA 字段时回退为 MHA vocab_sizehf_config.vocab_size, head_dimgetattr( hf_config, head_dim, hf_config.hidden_size // hf_config.num_attention_heads, ), hidden_dimhf_config.intermediate_size, norm_epshf_config.rms_norm_eps, rope_thetacls._get_rope_theta(hf_config, default10000.0), max_seq_lengetattr(hf_config, max_position_embeddings, 32768), attention_biasgetattr(hf_config, attention_bias, True), eos_idgetattr(hf_config, eos_token_id, 151645), enable_weight_tyinggetattr(hf_config, tie_word_embeddings, False), is_criticis_critic, attn_typekwargs.get(attn_type, BaseModelArgs.attn_type), )而 MoE 的判定模式qwen3展示了如何安全探测变体字段num_experts getattr(hf_config, num_experts, None) if num_experts is None: num_experts getattr(hf_config, num_local_experts, None) moe_enabled num_experts is not None and num_experts 1见 qwen3/model/args.py#L62-L67。原文档强调的 Critical 点所有字段映射必须逐项对照 HF 模型的config.json核实。此处映射错误不会立即报错而是造成下游静默错误。Step 4实现model.py—— 模型主体与基类契约model.py承载注意力、FFN、TransformerBlock 与顶层模型。需要适配的关键组件归一化RMSNorm或类似实现检查elementwise_affine是否可配置检查 epsilon 默认值若目标模型用LayerNorm则相应实现。Attention 模块Q/K/V 投影的 biasnn.Linear(..., biasTrue/False)QK norm有则加q_norm/k_norm无则删GQAn_kv_heads n_heads即为分组查询注意力Ulysses SP保留参考实现中的set_cp_group/_sp_enabled模式输出投影的 bias 存在性。FeedForward 模块SwiGLUw2(silu(w1(x)) * w3(x))现代 LLM 最常见检查各线性层 biasMoE 模型中指定层用MoE模块替换FeedForward。TransformerBlockpre-norm多数现代 LLM或 post-norm有 MoE 时用_is_moe_layer()检测 MoE 层。顶层模型ModelModel(BaseArchonModel)成员包括tok_embeddings、layers以ModuleDict组织、norm、output/scoreinit_weights()与 HF 的初始化方案保持一致init_buffers()RoPE 缓存 MoE buffersforward()必须遵循BaseArchonModel的签名。基类契约见 base.py#L144-L166class BaseArchonModel(nn.Module, ABC): abstractmethod def forward( self, tokens: torch.Tensor, positions: torch.Tensor, cu_seqlens: torch.Tensor, max_seqlen: int, tree_attn_meta: TreeAttentionMeta | None None, ) - torch.Tensor: ... abstractmethod def init_weights(self) - None: ... abstractmethod def init_buffers(self, buffer_device: torch.device | str) - None: ...注意forward采用打包packed序列接口cu_seqlens/max_seqlen描述变长序列的累积边界tree_attn_meta预留树注意力元信息。这意味着你的forward需要走 varlen 注意力路径而非逐样本填充。Step 5实现rope.py—— 旋转位置编码变体两条路线1标准 RoPE与 qwen2/qwen3 相同直接从 qwen2 重新导出from areal.experimental.models.archon.qwen2.model.rope import ( apply_rotary_emb, precompute_rope_cache, repeat_kv, reshape_for_broadcast, rotate_half, )2自定义 RoPEYaRN、NTK-aware 等自行实现precompute_rope_cache()与apply_rotary_emb()。核心差异通常只在inv_freq的计算方式缩放因子、插值等。Step 6实现state_dict_adapter.py—— 最容易出错的环节该适配器负责 HuggingFace 与 Archon 权重键名的双向映射。原文档明确将其标记为最易出错的步骤因为它直接决定权重能否完整、正确地加载与保存。必须处理的四类内容键名映射from_hf_map字典典型条目Embeddingmodel.embed_tokens.weight→tok_embeddings.weightAttentionmodel.layers.{}.self_attn.q_proj.weight→layers.{}.attention.wq.weightFFNmodel.layers.{}.mlp.gate_proj.weight→layers.{}.feed_forward.w1.weightNormmodel.layers.{}.input_layernorm.weight→layers.{}.attention_norm.weight输出头lm_head.weight→output.weight跳过项映射为Nonerotary_emb.inv_freq运行时计算模型特有键bias 项、QK norm 权重等。反向映射to_hf_map由from_hf_map自动生成MoE 专家权重如适用专家权重的 3D↔2D 转换可从 qwen3 复制 MoE 处理逻辑权重共享tie_word_embeddingsTrue时to_hf()需跳过output.weight。以 Qwen2StateDictAdapter 的真实实现对照可以看到上述每一条的具体落地from_hf_map中rotary_emb.inv_freq显式映射为Noneto_hf_map用循环从from_hf_map反向构建enable_weight_tying取自getattr(model_config, tie_word_embeddings, False)。此外还有两个容易被忽略的工程细节源码可查证weight tying 的加载侧补偿from_hf()在启用 tying 且 HF state dict 缺少lm_head.weight时用model.embed_tokens.weight补齐state_dict_adapter.py#L69-L77包裹前缀剥离convert_single_to_hf()需要剥掉激活检查点包装前缀._checkpoint_wrapped_module与torch.compile前缀._orig_mod否则检查点保存时键名对不上state_dict_adapter.py#L87-L102。基类提供的能力BaseStateDictAdapter构造函数接受model_config与可选的hf_assets_path后者若包含model.safetensors.index.json会解析出fqn_to_index_mapping支持多分片检查点的保存_load_safetensors_indexget_hf_storage_reader()返回HuggingFaceStorageReader供 DCP 直接读取 HF 检查点三个抽象方法构成子类契约from_hf()、to_hf()、convert_single_to_hf(name, tensor) - list[tuple[str, Tensor]]。验证方法原文档给出的 roundtrip 不变式# Roundtrip: archon - hf - archon 应保留全部键 hf_sd adapter.to_hf(archon_sd) roundtrip_sd adapter.from_hf(hf_sd) assert set(roundtrip_sd.keys()) set(archon_sd.keys())基类契约签名class ModelStateDictAdapter(BaseStateDictAdapter): def from_hf(self, hf_state_dict) - dict[str, Any]: ... def to_hf(self, archon_state_dict) - dict[str, Any]: ... def convert_single_to_hf(self, name, tensor) - list[tuple[str, torch.Tensor]]: ...Step 7实现parallelize.py—— 并行策略parallelize_model定义模型的并行方案。其函数签名必须满足 ParallelizeFn 协议def parallelize_model( model: nn.Module, parallel_dims: ArchonParallelDims, param_dtype: torch.dtype torch.bfloat16, reduce_dtype: torch.dtype torch.float32, loss_parallel: bool True, cpu_offload: bool False, reshard_after_forward_policy: str default, ac_config: ActivationCheckpointConfig | None None, enable_compile: bool True, ) - nn.Module:该协议在 model_spec.py 中以Protocol声明注释明确说明函数接收parallel_dims内部依据各*_enabled标志自行决定应用哪些并行策略。对照 qwen2 的 parallelize_qwen2实际签名与协议逐参数一致可直接作为模板。并行策略应用顺序原文档固定为如下顺序TPTensor Parallelism——跨设备切分 attention/FFNEPExpert Parallelism——仅 MoE 模型CPContext Parallelism / Ulysses SP——序列并行ACActivation Checkpointing——显存优化torch.compile——编译优化FSDPFully Sharded Data Parallelism——数据并行。按架构的关键适配点有 QK norm 的 Attentionwq/wk使用use_local_outputFalsenorm 需要 DTensor 输出并给q_norm/k_norm加SequenceParallel(sequence_dim2)无 QK norm 的 Attentionwq/wk/wv全部use_local_outputTrue带 bias 的 Attentionbias 项跟随其权重的同一并行计划MoE 层为 MoE 输入/输出、router gate、专家权重分别定义 TP plan从 qwen3 的apply_moe_ep_tp()与apply_non_moe_tp()复制纯 Dense 模型无 MoE 处理的简化计划从 qwen2 复制。qwen2 的实现还展示了该文件的典型工程构成导入apply_acactivation_checkpoint.py、apply_compilecompile.py、validate_cp_constraints/validate_tp_constraintsutils.py等公共工具说明新模型的 parallelize 函数应复用这些基础设施而非自行实现约束校验。Step 8编写spec.py并注册 ModelSpecModelSpec是把五块实现组装成一个整体规格的 dataclassmodel_spec.py#L85-L95name、model_class、model_args_class、state_dict_adapter_class、parallelize_fn、supported_model_types、pipelining_fn。标准模板对照 qwen2/spec.py 与 qwen3/spec.py 的真实写法from areal.experimental.models.archon.model_spec import ModelSpec, register_model_spec from areal.experimental.models.archon.pipeline_parallel import pipeline_llm from areal.experimental.models.archon.model.infra.parallelize import parallelize_model from areal.experimental.models.archon.model.model.args import ModelModelArgs from areal.experimental.models.archon.model.model.model import ModelModel from areal.experimental.models.archon.model.model.state_dict_adapter import ( ModelStateDictAdapter, ) MODEL_SPEC ModelSpec( nameModel, model_classModelModel, model_args_classModelModelArgs, state_dict_adapter_classModelStateDictAdapter, parallelize_fnparallelize_model, supported_model_typesfrozenset({model_type}), # 来自 HF config.json pipelining_fnpipeline_llm, ) # 模块被导入时自动注册 register_model_spec(MODEL_SPEC) __all__ [MODEL_SPEC]pipelining_fn使用 pipeline_parallel.py 中的pipeline_llm负责按 pipeline stage 切分模型从源码结构看其协议PipeliningFn返回 stages、model parts 以及当前 rank 是否持有首/尾 stage 的元组。注意supported_model_types应包含该实现处理的所有 HFmodel_type字符串。例如 qwen3 的实现同时覆盖 dense 与 MoE 两种形态因此注册为frozenset({qwen3, qwen3_moe})qwen3/spec.py#L18。漏掉变体字符串会导致引擎以该model_type查表时抛KeyError。Step 9在包级__init__.py中挂接自动注册在 areal/experimental/models/archon/init.py 中追加一行导入from areal.experimental.models.archon.model import spec as model_spec # noqa: F401该行触发模块导入即完成注册。当前仓库中的真实写法可作参照文件头注释特别说明了直接导入模块路径以避免先触发各模型包的__init__.pyfrom areal.experimental.models.archon.qwen2 import spec as qwen2_spec # noqa: F401 from areal.experimental.models.archon.qwen3 import spec as qwen3_spec # noqa: F401 from areal.experimental.models.archon.qwen3_5 import spec as qwen3_5_spec # noqa: F401忘记这一行是最常见的失误之一所有spec.py都写对了但引擎侧get_supported_model_types()枚举不到新类型。Step 10分层验证与测试原文档要求验证分阶段进行并先阅读现有测试再动手。仓库中 Archon 的测试位于 tests/experimental/archon/原文档给出的清单对应以下实际文件tests/experimental/archon/ conftest.py -- Pytest 配置版本检查 utils.py -- 共享工具模型加载、比较 test_qwen3_args.py -- Args 单元测试仅 CPU test_state_dict_adapter.py -- State dict 往返测试 test_weight_sync.py -- 权重完整性测试meta device test_forward.py -- 前向精度比较单 GPU test_hf_parity_qwen2.py -- 与 HuggingFace 的精度对齐 test_hf_parity_qwen3_moe.py -- MoE 变体的精度对齐 ...各阶段的测试写法以下代码模式引自原文档按模型复杂度裁剪Stage 1Args 测试仅 CPU必写用 mock HF config 验证from_hf_config()映射from unittest.mock import MagicMock def test_args_from_hf_config(): hf_config MagicMock() hf_config.hidden_size 4096 hf_config.num_hidden_layers 32 # ... 设置全部必需字段 args ModelModelArgs.from_hf_config(hf_config) assert args.dim 4096 assert args.n_layers 32Stage 2State Dict 适配器测试仅 CPU验证键映射往返def test_state_dict_roundtrip(): adapter ModelStateDictAdapter(mock_config) archon_sd {tok_embeddings.weight: torch.randn(vocab, dim), ...} hf_sd adapter.to_hf(archon_sd) roundtrip adapter.from_hf(hf_sd) assert set(roundtrip.keys()) set(archon_sd.keys())Stage 3权重完整性meta device仅 CPU验证模型每个参数都有对应的 HF 映射def test_weight_completeness(): with torch.device(meta): model ModelModel(args) adapter ModelStateDictAdapter(hf_config) for name, _ in model.named_parameters(): hf_pairs adapter.convert_single_to_hf(name, torch.empty(0)) assert len(hf_pairs) 0, fNo HF mapping for {name}Stage 4前向精度对齐单 GPU如可用对比 Archon 模型输出与 HuggingFace 参考实现pytest.mark.skipif(not torch.cuda.is_available(), reasonRequires CUDA) def test_forward_matches_hf(): # 分别加载 HF 与 Archon 模型 # 相同输入前向 # 在容差内比较 logits原文档的关键提醒不要硬编码测试分类。先检查 tests/experimental/archon/ 中现有测试文件如 test_qwen3_args.py、test_state_dict_adapter.py、test_weight_sync.py、test_forward.py遵循其 fixture 与 marker 约定并按模型特性裁剪测试范围——只有模型带 MoE 时才需要 MoE 专属测试。架构决策对照表原文档给出的决策地图直接回答目标模型的每个特征该抄谁特征qwen2qwen3目标模型中检查什么Attention bias有无HF config 的attention_biasQK norm无有HF config 的qk_norm或 modeling 文件中的 QKNorm 模块MoE无有HF config 的num_experts/num_local_expertsShared experts无有HF config 的num_shared_expertsDecoder sparse step无有HF config 的decoder_sparse_step权重共享两者均支持两者均支持HF config 的tie_word_embeddingsRoPE标准标准re-export qwen2HF modeling 代码中的 inv_freq 公式常见错误清单原文档 Common Mistakes 一节值得在完成后逐条自查未在state_dict_adapter.py中映射全部 HF 键导致权重静默丢失from_hf_config()字段映射写错用了错误的 HF config 属性名忘记处理from_hf_map中的None键如rotary_emb.inv_freq等应跳过的键模型带 MoE 时漏掉专家权重的 3D↔2D 转换有/无 QK norm 的 attention 用了错误的 TP planuse_local_output必须与 QK norm 的存在性匹配忘记在 areal/experimental/models/archon/init.py 添加导入行supported_model_typesfrozenset 未包含全部model_type变体使用print而非 areal.utils.logging 的getLogger()。完成检查清单收尾时按下表逐项确认文件存在且相互一致路径以仓库根目录为基准areal/experimental/models/archon/model/__init__.pyareal/experimental/models/archon/model/spec.py—— ModelSpec 定义 注册areal/experimental/models/archon/model/model/args.py—— ModelArgs from_hf_configareal/experimental/models/archon/model/model/model.py—— Model Attention FFNareal/experimental/models/archon/model/model/rope.py—— RoPE或 re-exportareal/experimental/models/archon/model/model/state_dict_adapter.py—— 键映射areal/experimental/models/archon/model/infra/parallelize.py—— 并行策略areal/experimental/models/archon/__init__.py—— 已添加导入行tests/experimental/archon/test_model_*.py—— 测试小结Archon 引擎的模型扩展体系可以概括为一条清晰的契约链BaseModelArgs管配置映射BaseArchonModel管前向与权重初始化BaseStateDictAdapter管检查点互转ParallelizeFn/PipeliningFn协议管并行与流水线最终由ModelSpecregister_model_spec收口为以 HFmodel_type为键的注册表。新模型的接入因此不是从零写一个模型而是在 qwen2/qwen3 参考实现上做一次受约束的差异化适配——最大的风险集中在state_dict_adapter的键映射与parallelize的 TP plan 两处而这两处分别有 roundtrip 不变式测试与 tests/experimental/archon/ 中现成的分布式测试test_distributed_tp.py、test_distributed_ep.py等可以复用验证。【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考