sparkle/mlx_streaming/core/prefetch/patch.py
fiser_jun 4745f264b2
2026-08-04 14:34:00 +08:00

57 lines
3.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""模型 patch把原生 MoE 块替换成流式块(文件后端 / 常驻切片版),并按需挂跨层预取。"""
from mlx_streaming import config
from mlx_streaming.core.moe.block import FileStreamingMoeBlock, StreamingMoeBlock
from mlx_streaming.core.prefetch.cross_layer import enable_cross_layer_prefetch
def patch_model_filebacked(model, store, hidden, moe_inter, group_size, bits,
proj_bits: dict | None = None,
layer_proj_bits: dict | None = None):
"""把每个 MoE 块替换为 FileStreamingMoeBlock并丢弃常驻的堆叠 switch_mlp。
storeFileExpertStore所有 MoE 层共用,按 (layer,expert) 缓存)。
proj_bits非空时走混合精度逐 proj 不同 bit专家须为对应混合重量化产出。
layer_proj_bits{绝对层号: {proj:bits}},非空时逐层用各自 proj_bits优先于 proj_bits
对应 requantize_dir_layered 产出。各层 QSL 用该层 bit与流式存盘文件一一对应。
返回被替换的层数。被替换后原 switch_mlp 不再被引用,惰性权重不会被物化。
"""
patched = 0
for i, layer in enumerate(model.layers):
layer._layer_idx = i
# 用 object.__setattr__ 存反向 model 引用,避免进 nn.Module(=dict)的子模块树:
# 否则 model→layer→model 成环,wired_limit 等 tree_reduce 遍历会无限递归(RecursionError)。
object.__setattr__(layer, "_prefetch_model_ref", model)
mlp = getattr(layer, "mlp", None)
if mlp is not None and hasattr(mlp, "switch_mlp") and hasattr(mlp, "gate"):
pb = layer_proj_bits.get(i, proj_bits) if layer_proj_bits else proj_bits
# 捕获共享专家引用使其常驻Qwen3-Next 有Qwen3-MoE 无)
layer.mlp = FileStreamingMoeBlock(
gate=mlp.gate, top_k=mlp.top_k, norm_topk_prob=mlp.norm_topk_prob,
store=store, layer_idx=i, hidden=hidden, moe_inter=moe_inter,
group_size=group_size, bits=bits, proj_bits=pb,
shared_expert=getattr(mlp, "shared_expert", None),
shared_expert_gate=getattr(mlp, "shared_expert_gate", None),
)
# 同理:不进 dict,避免 mlp→model→...→mlp 成环。供 native-fused-prefetch 取下层 gate。
object.__setattr__(layer.mlp, "_prefetch_model_ref", model)
patched += 1
if config.cross_layer_prefetch() or getattr(store, "_staging", None) is not None:
enable_cross_layer_prefetch()
return patched
def patch_model(model, store_factory=None):
"""把模型里每个 MoE 块(含 switch_mlp 与 gate 的块)替换成 StreamingMoeBlock。
store_factory(layer_idx)->LruExpertStore可为 None仅做 uniq 切片、不接磁盘后端)。
返回被替换的层数,便于校验确实命中了 MoE 层。
"""
patched = 0
for i, layer in enumerate(model.layers):
mlp = getattr(layer, "mlp", None)
if mlp is not None and hasattr(mlp, "switch_mlp") and hasattr(mlp, "gate"):
store = store_factory(i) if store_factory is not None else None
layer.mlp = StreamingMoeBlock(mlp, layer_idx=i, store=store)
patched += 1
return patched