diff --git a/configs/deepseek-v4-flash-dspark-moe.json b/configs/deepseek-v4-flash-dspark-moe.json new file mode 100644 index 000000000..4e7f9410a --- /dev/null +++ b/configs/deepseek-v4-flash-dspark-moe.json @@ -0,0 +1,60 @@ +{ + "architectures": ["DSparkDraftModel"], + "attention_bias": false, + "attention_dropout": 0.0, + "auto_map": {"AutoModel": "dspark.DSparkDraftModel"}, + "block_size": 7, + "bos_token_id": 0, + "dflash_config": { + "attention_mode": "gqa", + "confidence_head_alpha": 1.0, + "confidence_head_with_markov": true, + "enable_confidence_head": true, + "markov_head_type": "vanilla", + "markov_rank": 256, + "mask_token_id": 128799, + "moe_bias_update_rate": 0.001, + "moe_dispatch": "grouped_mm", + "projector_type": "dspark", + "target_layer_ids": [1, 11, 21, 31, 41] + }, + "dtype": "bfloat16", + "eos_token_id": 1, + "head_dim": 128, + "hidden_act": "silu", + "hidden_size": 4096, + "initializer_range": 0.02, + "intermediate_size": 12288, + "layer_types": [ + "full_attention", + "full_attention", + "full_attention", + "full_attention", + "full_attention" + ], + "max_position_embeddings": 1048576, + "max_window_layers": 5, + "model_type": "qwen3", + "moe_intermediate_size": 2048, + "moe_preset": "deepseek_v4", + "n_routed_experts": 64, + "n_shared_experts": 1, + "num_attention_heads": 32, + "num_experts_per_tok": 6, + "num_hidden_layers": 5, + "num_key_value_heads": 8, + "num_target_layers": 43, + "pad_token_id": 1, + "rms_norm_eps": 1e-06, + "rope_parameters": { + "factor": 16.0, + "original_max_position_embeddings": 65536, + "rope_theta": 10000.0, + "rope_type": "yarn" + }, + "sliding_window": null, + "tie_word_embeddings": false, + "use_cache": true, + "use_sliding_window": false, + "vocab_size": 129280 +} diff --git a/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md b/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md index e62b4e013..8876e4485 100644 --- a/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md +++ b/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md @@ -103,6 +103,30 @@ MI355X node over the full two epochs (1,885 optimizer steps): 6.32 s per 128-sample step on average, 6.1-6.4 s at steady state (about 3.1 s waiting for capture and 3.0 s of trainer compute), 3 h 18 min end to end. +## MoE-FFN arm (ablation) + +`examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-disaggregated.yaml` +is the same recipe with `configs/deepseek-v4-flash-dspark-moe.json`: the +five-layer GQA decoder keeps its attention, and each layer's dense MLP becomes +the target's MoE (`moe_preset: deepseek_v4`: sqrt-softplus scores, +aux-loss-free top-k with the sign-controlled selection bias, combine weights +renormalized and scaled by 1.5, one ungated shared expert, SwiGLU clamped at +10). Sizes are per run: 64 routed experts, top-6, width 2048, so the activated +FFN width (6 x 2048 + 2048 shared) matches the dense 12288 at ~10x the FFN +parameters. Run it against the dense recipe with identical hparams; the two +YAMLs differ only in the draft JSON and run names. The capture servers are +shared by both arms unchanged. + +Training-only knobs live under the draft JSON's `dflash_config`: +`moe_bias_update_rate` (0.001, the balancing controller's step) and +`moe_dispatch` (`grouped_mm` runs the experts as grouped GEMMs with no host +sync; `sorted_loop` is the portable fallback). The trainer logs `moe/*` load +metrics (max/min load ratios, unused-expert fraction, balancing-bias +magnitude) alongside the usual scalars. Checkpoints keep the official +per-expert naming (`layers.N.mlp.experts.{i}.w{1,2,3}.weight`, +`layers.N.mlp.gate.bias`, `layers.N.mlp.shared_experts.w{1,2,3}.weight`), so +exports load into SGLang's DeepSeek-V4 MoE unchanged. + ## Fresh attempts Delete the run's `outputs/` directory and, whenever a capture server was diff --git a/docs/sections/advanced_features/customization.md b/docs/sections/advanced_features/customization.md index c7b58f870..a2ae5cbc1 100644 --- a/docs/sections/advanced_features/customization.md +++ b/docs/sections/advanced_features/customization.md @@ -128,6 +128,42 @@ drafts currently implements the GQA/MHA layout only, so plan benchmarks accordingly. DFlash2 otherwise follows the same mode selection; its convolution and selector do not change the attention projection contract. +## MoE FFN for DFlash-family drafts + +Any DFlash-family draft (DFlash, DFlash2, DSpark) swaps its dense MLP for a +sparse MoE FFN when the draft JSON sets `n_routed_experts > 0`. The MoE is one +configurable layer (`specforge/modeling/draft/moe/`, see its `DESIGN.md`): +a `moe_preset` names a target family's routing recipe, and the architecture +keys use the target checkpoints' native HF names so they can be copied from +the target's `config.json`: + +```json +{ + "moe_preset": "deepseek_v4", + "n_routed_experts": 64, + "num_experts_per_tok": 6, + "moe_intermediate_size": 2048, + "n_shared_experts": 1, + "dflash_config": {"moe_bias_update_rate": 0.001, "moe_dispatch": "grouped_mm"} +} +``` + +Top-level keys override the preset (for ablations: `scoring_func`, +`norm_topk_prob`, `routed_scaling_factor`, `balance`, `shared_expert_gate`, +`swiglu_limit`, ...). Training-only knobs live under `dflash_config` with an +`moe_` prefix and never change the checkpoint. Checkpoints, warm starts and +exports keep the official per-expert naming (`experts.{i}.w{1,2,3}.weight`), +so an exported drafter loads into SGLang unchanged. Dense drafts are +unaffected: with no `n_routed_experts` the kernel provider's MLP is used as-is. + +`deepseek_v4` is the checked-in preset (DeepSeek-V4 routing: +`sqrtsoftplus` scores, aux-loss-free `noaux_tc` balancing, combine weights +renormalized and scaled by 1.5, one ungated shared expert, SwiGLU clamp 10); +`configs/deepseek-v4-flash-dspark-moe.json` uses it. A new target family is a +preset registration plus whichever components it needs (score function, +balance controller, experts backend, shared expert); each registers by name +from its own module. + ## Draft architectures Draft classes register through `@register_draft`. The key defaults to the diff --git a/examples/configs/README.md b/examples/configs/README.md index 905c4eb10..18b6f47a7 100644 --- a/examples/configs/README.md +++ b/examples/configs/README.md @@ -88,6 +88,11 @@ servers. Its [runbook](../../docs/recipes/deepseek-v4-flash-dspark-disaggregated.md) covers the v0.5.18 SGLang capture patch and the bundled `deepseek-v4` chat template (the checkpoint ships no Jinja template). +`deepseek-v4-flash-dspark-moe-disaggregated.yaml` is the MoE-FFN arm of the +drafter-architecture ablation: the same recipe with +`configs/deepseek-v4-flash-dspark-moe.json`, whose `moe_preset: deepseek_v4` +swaps the dense MLP for the target's routing (64 routed + 1 shared experts, +top-6, width 2048); see the runbook's MoE section. `qwen3.8-27b-dflash2-disaggregated.yaml` (external services, two nodes) and its managed-local siblings `qwen3.8-27b-dflash2-4server-dp4-disaggregated.yaml` diff --git a/examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-disaggregated.yaml b/examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-disaggregated.yaml new file mode 100644 index 000000000..615db8755 --- /dev/null +++ b/examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-disaggregated.yaml @@ -0,0 +1,97 @@ +# MoE-FFN arm of the DSpark drafter-architecture ablation for +# DeepSeek-V4-Flash. Identical to deepseek-v4-flash-dspark-disaggregated.yaml +# (the dense arm) except the draft JSON and the run/store names, so the A/B +# diff is the FFN: 64 routed experts + 1 shared, top-6, moe_intermediate 2048 +# (activated width 6x2048 + 2048 ~= the dense 12288) with DeepSeek-V4 routing. +# If a from-scratch MoE run destabilizes at the shared hparams, lower the LR +# and clip on BOTH arms rather than on this file alone. +model: + target_model_path: deepseek-ai/DeepSeek-V4-Flash-0731 + draft_model_config: configs/deepseek-v4-flash-dspark-moe.json + target_backend: sglang + trust_remote_code: true + # DeepSeek-V4 checkpoints use the native inference weight layout. + embedding_key: embed.weight + lm_head_key: head.weight + mask_token_id: 128799 + torch_dtype: bfloat16 + sglang_mem_fraction_static: 0.85 + # data.max_length plus headroom: capture rejects inputs at exactly context length. + sglang_context_length: 8704 + sglang_max_running_requests: 8 + # Routed experts are fp4; the default MoE path cannot run them on B200. + sglang_moe_runner_backend: flashinfer_mxfp4 + +data: + train_data_path: ./cache/dataset/sharegpt_train.jsonl + max_length: 8192 + chat_template: deepseek-v4 + cache_dir: cache + build_dataset_num_proc: 64 + # Async prefetch: overlaps per-sample Mooncake feature fetches with compute. + dataloader_num_workers: 8 + +training: + strategy: dspark + num_epochs: 2 + # 4 ranks x 32 microbatches -> global batch 128. + batch_size: 1 + accumulation_steps: 32 + learning_rate: 0.0006 + lr_scheduler: constant + warmup_ratio: 0 + max_grad_norm: 1 + attention_backend: flex_attention + num_anchors: 512 + loss_decay_gamma: 4.0 + objective_chunk_blocks: 128 + dspark_ce_loss_alpha: 0.1 + dspark_l1_loss_alpha: 0.9 + dspark_confidence_head_alpha: 1.0 + save_interval: 128 + log_interval: 10 + dist_timeout: 30 + seed: 42 + prompt_seed: 1 + +tracking: + report_to: wandb + wandb_project: specforge + wandb_name: deepseek-v4-flash-dspark-moe-disaggregated + wandb_dir: outputs/deepseek-v4-flash-dspark-moe-disaggregated/wandb + +runtime: + producer_lease: 8 + producer_concurrency: 8 + # Keep two 128-sample optimizer quanta in flight (~0.4 GiB features/sample); + # the low watermark is one full quantum so complete windows always dispatch. + in_flight_high_watermark: 256 + in_flight_low_watermark: 128 + resident_high_watermark_bytes: 137438953472 + resident_low_watermark_bytes: 103079215104 + feature_store_max_resident_bytes: 171798691840 + +run_id: deepseek-v4-flash-dspark-moe-disaggregated +output_dir: outputs/deepseek-v4-flash-dspark-moe-disaggregated + +deployment: + mode: disaggregated + trainer: + nnodes: 1 + nproc_per_node: 4 + disaggregated: + control_dir: outputs/deepseek-v4-flash-dspark-moe-disaggregated/control + consumer_state_dir: outputs/deepseek-v4-flash-dspark-moe-disaggregated/consumer-state + backend: mooncake + store_id: deepseek-v4-flash-dspark-moe-disaggregated + # Two TP2 servers out-produce one TP4 server (TP prefill scaling is sublinear). + server_urls: + - http://127.0.0.1:30000 + - http://127.0.0.1:30001 + mooncake_metadata_server: http://127.0.0.1:35880/metadata + mooncake_master_server_addr: 127.0.0.1:35551 + mooncake_local_hostname: 127.0.0.1 + mooncake_protocol: tcp + client_buffer_size: 1073741824 + idle_timeout_s: 7200 + peer_wait_timeout_s: 7200 diff --git a/scripts/warm_start_moe_drafter.py b/scripts/warm_start_moe_drafter.py new file mode 100644 index 000000000..449f9525b --- /dev/null +++ b/scripts/warm_start_moe_drafter.py @@ -0,0 +1,112 @@ +#!/usr/bin/env python3 +"""Build a warm-start source for an MoE DFlash-family drafter from a DeepSeek-V4 target. + +Constructs the draft model from its config (random init for attention, heads, +norms, projections), seeds every MoE layer's routed experts, gate weight, +noaux bias and shared expert from one target layer (dequantized fp4/fp8 -> +bf16), and writes an HF-format directory usable as +``model.draft_checkpoint_path``. Requires the draft's expert shape to match the +target's (256 x 2048 for DeepSeek-V4-Flash); a smaller draft takes a strided +subset of experts. + +Example: + python scripts/warm_start_moe_drafter.py \ + --draft-config examples/configs/kan-ablations/deepseek-v4-flash-dspark-moe256-auxbal.json \ + --target-snapshot /cluster-storage/models/models--deepseek-ai--DeepSeek-V4-Flash-0731/snapshots/ \ + --target-layers 3,11,21,31,41 --output-dir warm-starts/moe256-from-0731 +""" + +from __future__ import annotations + +import argparse +import json +import os +import time + +import torch + + +def main() -> int: + ap = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter + ) + ap.add_argument("--draft-config", required=True) + ap.add_argument( + "--target-snapshot", required=True, help="local HF snapshot dir of the target" + ) + ap.add_argument( + "--target-layers", + required=True, + help="comma list, one target layer per draft layer", + ) + ap.add_argument("--output-dir", required=True) + ap.add_argument("--seed", type=int, default=42) + ap.add_argument( + "--select", + default="strided", + help="expert selection when draft has fewer experts", + ) + args = ap.parse_args() + + from specforge.modeling.auto import AutoDraftModel, AutoDraftModelConfig + from specforge.modeling.draft.moe import ( + apply_warm_start, + iter_moe_layers, + plan_warm_start, + resolve_moe_config, + to_checkpoint_state_dict, + ) + from specforge.modeling.draft.moe.deepseek_v4_target import load_target_moe_layer + + target_layers = [int(x) for x in args.target_layers.split(",")] + config = AutoDraftModelConfig.from_file(args.draft_config) + moe_cfg = resolve_moe_config(config) + if moe_cfg is None: + raise SystemExit("draft config is dense (n_routed_experts == 0)") + torch.manual_seed(args.seed) + t0 = time.time() + model = AutoDraftModel.from_config(config, torch_dtype=torch.bfloat16) + layers = list(iter_moe_layers(model)) + if len(layers) != len(target_layers): + raise SystemExit( + f"{len(layers)} MoE layers but {len(target_layers)} target layers given" + ) + print( + f"built draft ({sum(p.numel() for p in model.parameters())/1e9:.2f}B params) in {time.time()-t0:.0f}s" + ) + + target_cfg = json.load(open(os.path.join(args.target_snapshot, "config.json"))) + n_target = int(target_cfg["n_routed_experts"]) + for i, (layer, tl) in enumerate(zip(layers, target_layers)): + t1 = time.time() + source = load_target_moe_layer(args.target_snapshot, tl) + plan = plan_warm_start( + layer.cfg, n_target_experts=n_target, strategy=args.select + ) + loaded = apply_warm_start(layer, plan, source) + print( + f"draft layer {i} <- target layer {tl}: {len(loaded)} tensors, " + f"experts {plan.target_expert_ids[:3]}...{plan.target_expert_ids[-1]} ({time.time()-t1:.0f}s)" + ) + + for key, value in moe_cfg.serving_fields().items(): + setattr(model.config, key, value) + model.config.moe_warm_start = { + "target": os.path.basename( + os.path.dirname(os.path.dirname(args.target_snapshot.rstrip("/"))) + ), + "snapshot": os.path.basename(args.target_snapshot.rstrip("/")), + "target_layers": target_layers, + "select": args.select, + "seed": args.seed, + } + os.makedirs(args.output_dir, exist_ok=True) + model.save_pretrained( + args.output_dir, state_dict=to_checkpoint_state_dict(model.state_dict()) + ) + print(f"wrote {args.output_dir} in {time.time()-t0:.0f}s total") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/specforge/export/checkpoint_io.py b/specforge/export/checkpoint_io.py index 87aad6929..97c9347aa 100644 --- a/specforge/export/checkpoint_io.py +++ b/specforge/export/checkpoint_io.py @@ -159,7 +159,12 @@ def materialize_draft( draft_config = AutoDraftModelConfig.from_file(draft_config_path) model = AutoDraftModel.from_config(draft_config, torch_dtype=torch.bfloat16) - missing, unexpected = model.load_state_dict(state["draft_state_dict"], strict=False) + from specforge.modeling.draft.moe import from_checkpoint_state_dict + + # Files use the official naming; modules may use a native MoE layout. + missing, unexpected = model.load_state_dict( + from_checkpoint_state_dict(state["draft_state_dict"]), strict=False + ) if unexpected: raise ValueError( f"checkpoint carries weights the {type(model).__name__} architecture " diff --git a/specforge/export/to_hf.py b/specforge/export/to_hf.py index 928408b14..8ed702321 100644 --- a/specforge/export/to_hf.py +++ b/specforge/export/to_hf.py @@ -31,6 +31,7 @@ materialize_draft, resolve_training_state, ) +from specforge.modeling.draft.moe import resolve_moe_config, to_checkpoint_state_dict def _load_embedding_tensor(source: str, key: str) -> torch.Tensor: @@ -87,7 +88,13 @@ def export_to_hf( model = materialize_draft( state, draft_config_path, vocab_mapping_path=vocab_mapping_path ) - full_state = dict(model.state_dict()) + full_state = dict(to_checkpoint_state_dict(model.state_dict())) + moe_cfg = resolve_moe_config(model.config) + if moe_cfg is not None: + # Serving engines read the routing recipe from config.json; the draft + # JSON only names a preset, so materialize the resolved fields. + for key, value in moe_cfg.serving_fields().items(): + setattr(model.config, key, value) owns_embedding = hasattr(model, "embed_tokens") if owns_embedding and "embed_tokens.weight" not in state["draft_state_dict"]: if not embedding_source: diff --git a/specforge/export/to_sglang.py b/specforge/export/to_sglang.py index b08b856ec..88a4f2449 100644 --- a/specforge/export/to_sglang.py +++ b/specforge/export/to_sglang.py @@ -28,6 +28,7 @@ materialize_draft, resolve_training_state, ) +from specforge.modeling.draft.moe import to_checkpoint_state_dict #: per-architecture trainer-key -> serving-key renames ({} = identity). WEIGHT_MAPS: Dict[str, Dict[str, str]] = { @@ -82,7 +83,11 @@ def export_to_sglang( weight_map = WEIGHT_MAPS.get(type(model).__name__, {}) # the model's state dict includes any refreshed t2d/d2t buffers; drop the # embeddings exactly as the trainer-side checkpoint filter does. - full = {k: v for k, v in model.state_dict().items() if "embed" not in k.lower()} + full = { + k: v + for k, v in to_checkpoint_state_dict(model.state_dict()).items() + if "embed" not in k.lower() + } model.save_pretrained(output_dir, state_dict=_serving_state(full, weight_map)) apply_legacy_rope_scaling(output_dir) return output_dir diff --git a/specforge/modeling/auto.py b/specforge/modeling/auto.py index de19dc5c7..8b52a0f45 100644 --- a/specforge/modeling/auto.py +++ b/specforge/modeling/auto.py @@ -62,6 +62,32 @@ def filtered_warning(msg): config = AutoConfig.from_pretrained(pretrained_model_name_or_path) model_cls = cls._model_cls_from_config(config) kwargs = {**kwargs, "config": config} + state_dict = load_native_state_dict(config, pretrained_model_name_or_path) + if state_dict is not None: + # HF from_pretrained assigns tensors by key and refuses an + # explicit state_dict alongside a path, so build the module and + # load the converted state ourselves. + torch_dtype = kwargs.pop("torch_dtype", kwargs.pop("dtype", None)) + output_loading_info = bool(kwargs.pop("output_loading_info", False)) + model = model_cls._from_config(config, torch_dtype=torch_dtype) + result = model.load_state_dict(state_dict, strict=False) + del state_dict + missing = [k for k in result.missing_keys if "embed_tokens" not in k] + if missing or result.unexpected_keys: + raise ValueError( + f"{pretrained_model_name_or_path!r} does not match " + f"{model_cls.__name__}: missing {missing[:5]}, " + f"unexpected {list(result.unexpected_keys)[:5]}" + ) + model.eval() + if output_loading_info: + return model, { + "missing_keys": list(result.missing_keys), + "unexpected_keys": list(result.unexpected_keys), + "mismatched_keys": [], + "error_msgs": [], + } + return model model = model_cls.from_pretrained( pretrained_model_name_or_path, *model_args, **kwargs ) @@ -71,6 +97,41 @@ def filtered_warning(msg): return model +def load_native_state_dict(config, pretrained_model_name_or_path): + """Checkpoint files use the official parameter naming; modules may use a + native layout (MoE experts). HF ``from_pretrained`` assigns tensors by key + and cannot regroup them, so read the files (memory-mapped) and convert at + this boundary. Returns ``None`` when no conversion is needed (dense + drafts). Callers that only need the tensors (warm start) use this directly + instead of materializing a second model.""" + from specforge.modeling.draft.moe import ( + from_checkpoint_state_dict, + is_moe_config, + ) + + if not is_moe_config(config): + return None + import glob + + from safetensors import safe_open + + path = str(pretrained_model_name_or_path) + if not os.path.isdir(path): + from huggingface_hub import snapshot_download + + path = snapshot_download(path, allow_patterns=["*.safetensors", "*.json"]) + files = sorted(glob.glob(os.path.join(path, "*.safetensors"))) + if not files: + raise FileNotFoundError(f"no safetensors weights under {path!r}") + state = {} + for file in files: + # mmap-backed views: only the regrouped tensors are materialized. + with safe_open(file, framework="pt", device="cpu") as handle: + for key in handle.keys(): + state[key] = handle.get_tensor(key) + return from_checkpoint_state_dict(state) + + class AutoDraftModelConfig: @classmethod def from_file(cls, config_path: str): diff --git a/specforge/modeling/draft/dflash.py b/specforge/modeling/draft/dflash.py index ea50ec122..68fe9f6ce 100644 --- a/specforge/modeling/draft/dflash.py +++ b/specforge/modeling/draft/dflash.py @@ -21,6 +21,7 @@ from typing_extensions import Tuple, Unpack from .dflash_kernels import DEFAULT_DFLASH_KERNELS, DFlashKernels +from .moe import MoELayer, apply_pending_balance_updates, build_ffn from .flex_attention_backend import flex_attention_backend from .registry import register_draft @@ -558,7 +559,9 @@ def __init__( layer_idx=layer_idx, kernels=kernels, ) - self.mlp = kernels.make_mlp(config) + # Dense MLP from the kernel provider, or an MoELayer when the draft + # JSON sets n_routed_experts > 0 (see modeling/draft/moe). + self.mlp = build_ffn(config, dense=kernels.make_mlp) self.input_layernorm = kernels.make_rms_norm( config.hidden_size, config.rms_norm_eps ) @@ -733,6 +736,14 @@ def _build_decoder_layer( return self.decoder_layer_class(config, layer_idx, kernels) + def _init_weights(self, module: nn.Module) -> None: + if isinstance(module, MoELayer): + # MoE components hold bare Parameters (gate weight, stacked + # experts) that the inherited Qwen3 init never visits. + module.reset_parameters(std=self.config.initializer_range) + return + super()._init_weights(module) + def _init_draft_head(self, config, dflash_config: dict) -> None: del config, dflash_config @@ -796,6 +807,11 @@ def forward( **kwargs, ) -> CausalLMOutputWithPast: hidden_states = noise_embedding + if self.training: + # Consume the PREVIOUS forward's routing statistics before any + # routing this step, outside every activation-checkpoint region + # (see modeling/draft/moe/balance.py for why the timing matters). + apply_pending_balance_updates(self) target_hidden = self.hidden_norm(self.fc(target_hidden)) position_embeddings = self.rotary_emb(hidden_states, position_ids) for layer_type, layer in zip(self.layer_types, self.layers): diff --git a/specforge/modeling/draft/moe/DESIGN.md b/specforge/modeling/draft/moe/DESIGN.md new file mode 100644 index 000000000..15bd32538 --- /dev/null +++ b/specforge/modeling/draft/moe/DESIGN.md @@ -0,0 +1,112 @@ +# MoE FFN Design (`specforge.modeling.draft.moe`) + +Design note for the sparse-MoE FFN that any DFlash-family draft (DFlash, +DFlash2, DSpark) can opt into. The training plane's picture is in +[`../../../training/DESIGN.md`](../../../training/DESIGN.md). + +## Responsibility + +Owns the FFN of a decoder layer when the draft JSON sets +`n_routed_experts > 0`, and nothing else: attention, heads, losses and the +trainer loop are unchanged. The package is **one configurable layer**, not one +block per target family. This is the Megatron-Core split (one `MoELayer`, +router/balancing/experts/dispatcher/shared-expert as orthogonal components +picked by config), chosen over the transformers pattern of a copied +`XxxSparseMoeBlock` per model because the family differences are a handful of +small functions while the expensive parts are shared: + +| differs per target family | shared by every family | +| -------------------------------- | ----------------------------------------- | +| score function (softmax, sigmoid, sqrtsoftplus) | token dispatch (sorted segments, grouped GEMM) | +| balancing policy (aux loss, aux-loss-free bias, none) | deferred balance-update timing vs activation checkpointing | +| combine-weight renorm and scale | FSDP-friendly stacked expert layout | +| shared-expert gate (none, sigmoid) | checkpoint naming boundary | +| SwiGLU clamp | load metrics, warm-start plans | + +## Why match the target's MoE + +The drafter's MoE should be the *target's* MoE, expressed as a preset: + +- **Warm start.** Same expert shape means the draft experts can be seeded from + a subset of the target's, which a dense drafter cannot do. +- **Serving.** The drafter runs inside SGLang's draft model; matching the + target's routing reuses its fused MoE kernels and weight naming. +- **Latency.** At small batch, top-k of narrow experts reads about the same + bytes as the dense MLP, so MoE buys parameters at roughly constant step + cost. That is the hypothesis the dense-vs-MoE ablation tests. + +## Layout + +``` +config.py MoEConfig + preset registry. Architecture keys use the target + checkpoints' native HF names at the draft JSON top level; + training-only knobs live under dflash_config as moe_*. +router.py Router contract (x -> RoutingResult) + score-function registry. +balance.py BalanceController contract; owns selection-bias buffers, stashes + counts in forward, applies updates from the model forward. +experts.py RoutedExperts contract (weights + dispatch), MoEConfig.dispatch knob. +shared.py SharedExpert contract, gate variant via MoEConfig.shared_expert_gate. +layer.py MoELayer = gate + experts + shared_experts; build_ffn() is the + dense/MoE switch used by the DFlash decoder layer. +hooks.py apply_pending_balance_updates / collect_moe_aux_loss / + collect_moe_metrics over any module tree. +state_dict.py to/from_checkpoint_state_dict: module layout <-> official names. +init.py WarmStartPlan: which target experts seed which draft experts. +``` + +Implementations register into these registries at import time (imported at +the bottom of `__init__.py`): + +``` +topk_router.py "topk" router; score functions softmax / sigmoid / sqrtsoftplus; + optional group-limited selection (DeepSeek top-2 group scores). +noaux_tc.py "noaux_tc" controller: fp32 selection bias + sign controller on + all-reduced loads; converter gate.balance.bias <-> gate.bias. +grouped_experts.py "grouped" experts: stacked [E, out, in] w1/w2/w3, sorted-segment + loop or torch._grouped_mm dispatch; converter experts.w1 <-> + experts.{i}.w1.weight. +swiglu_shared.py "swiglu" ungated shared expert (shared_experts.w1/w2/w3). +presets.py "deepseek_v4": sqrtsoftplus + noaux_tc + renorm x1.5 + one + ungated shared expert + SwiGLU clamp 10. +``` + +## Contracts that matter + +**Attribute names are the checkpoint contract.** `MoELayer` exposes `gate`, +`experts`, `shared_experts` (the DeepSeek-family names SGLang loads). A +component whose native parameter layout differs from the official file naming +registers a `state_dict` converter pair; both directions are idempotent and +no-ops on dense models. + +**Naming is converted at the boundary, not in `state_dict()`.** FSDP's +full-state-dict hooks index the gathered dict by the module's own parameter +FQNs, so a rename inside `state_dict()` breaks under `use_orig_params`. Every +file read/write goes through `to_checkpoint_state_dict` / +`from_checkpoint_state_dict`: `FSDPTrainingBackend` save/load, warm start, +`materialize_draft`, and the HF and SGLang exporters. + +**Balance updates are deferred.** `BalanceController.observe` only stashes +(overwrite, never accumulate) so an activation-checkpoint recompute leaves +identical state. `apply_pending_update` runs from the *model* forward before +any routing, outside checkpoint regions, and may run collectives. Mutating +selection state inside a layer forward would make the recompute route +differently and raise `CheckpointError`. + +**Metrics ride the existing scalar channel.** `collect_moe_metrics` yields +`moe/load_max_ratio`, `moe/load_min_ratio`, `moe/experts_unused_frac` plus +controller metrics; the DFlash/DSpark strategies add them to `StepOutput.metrics`, +and the trainer DP-averages and logs them like any other scalar. + +**Aux losses are collected, not yet consumed.** `collect_moe_aux_loss` sums +scaled layer losses; wiring it into an objective is done with the first preset +whose balancing policy emits one (aux-loss-free policies do not). + +## Extension points + +- New target family: `register_moe_preset("", scoring_func=..., ...)` + plus any missing component registrations. The draft JSON then sets + `moe_preset` and the per-run sizes. +- New dispatch: a `RoutedExperts` subclass under `register_experts_backend`, + or a new `MoEConfig.dispatch` value handled inside an existing backend. +- Ablation knobs: any `ARCHITECTURE_KEYS` entry at the draft JSON top level + overrides its preset default; `dflash_config.moe_*` overrides training knobs. diff --git a/specforge/modeling/draft/moe/__init__.py b/specforge/modeling/draft/moe/__init__.py new file mode 100644 index 000000000..3cd8ad691 --- /dev/null +++ b/specforge/modeling/draft/moe/__init__.py @@ -0,0 +1,137 @@ +# coding=utf-8 +# Copyright 2024 The SpecForge team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +"""Sparse MoE FFN for DFlash-family drafts. + +One configurable layer, not one block per target family. A draft JSON selects a +``moe_preset`` (the routing recipe of a target family) and may override any +architecture knob for ablations; the expensive parts (dispatch, checkpoint +naming, FSDP layout, deferred balance updates) are shared. See ``DESIGN.md``. + +Modules: + +- :mod:`.config` ``MoEConfig`` + the preset registry; resolves a draft JSON +- :mod:`.router` ``Router`` contract (scores -> top-k) + score functions +- :mod:`.balance` ``BalanceController`` contract: balancing as a policy +- :mod:`.experts` ``RoutedExperts`` contract: expert weights + dispatch +- :mod:`.shared` ``SharedExpert`` contract +- :mod:`.layer` ``MoELayer`` composition; ``build_ffn`` is the dense/MoE switch +- :mod:`.hooks` model-level plumbing: balance updates, aux loss, metrics +- :mod:`.state_dict` module layout <-> official checkpoint naming boundary +- :mod:`.init` warm-start plans from a target model's experts + +Implementations register into the registries from their own modules +(:mod:`.topk_router`, :mod:`.noaux_tc`, :mod:`.grouped_experts`, +:mod:`.swiglu_shared`) and presets in :mod:`.presets`; the modules above +define contracts and hold no routing math themselves. +""" + +# Implementations and presets register at import time. +from . import ( # noqa: E402,F401 isort: skip + grouped_experts, + noaux_tc, + presets, + swiglu_shared, + topk_router, +) +from .balance import ( + BALANCE_CONTROLLERS, + BalanceController, + build_balance_controller, + register_balance_controller, +) +from .config import ( + MOE_PRESETS, + MoEConfig, + available_moe_presets, + is_moe_config, + register_moe_preset, + resolve_moe_config, +) +from .experts import ( + EXPERTS_BACKENDS, + RoutedExperts, + build_routed_experts, + register_experts_backend, +) +from .hooks import ( + apply_pending_balance_updates, + collect_moe_aux_loss, + collect_moe_metrics, + iter_moe_layers, +) +from .init import ( + WarmStartPlan, + apply_warm_start, + plan_warm_start, + select_target_experts, +) +from .layer import MoELayer, build_ffn +from .router import ( + ROUTERS, + SCORE_FUNCTIONS, + Router, + RoutingResult, + build_router, + get_score_function, + register_router, + register_score_function, +) +from .shared import ( + SHARED_EXPERTS, + SharedExpert, + build_shared_expert, + register_shared_expert, +) +from .state_dict import ( + from_checkpoint_state_dict, + register_state_dict_converter, + to_checkpoint_state_dict, +) + +__all__ = [ + "BALANCE_CONTROLLERS", + "BalanceController", + "EXPERTS_BACKENDS", + "MOE_PRESETS", + "MoEConfig", + "MoELayer", + "ROUTERS", + "RoutedExperts", + "Router", + "RoutingResult", + "SCORE_FUNCTIONS", + "SHARED_EXPERTS", + "SharedExpert", + "WarmStartPlan", + "apply_pending_balance_updates", + "apply_warm_start", + "available_moe_presets", + "build_balance_controller", + "build_ffn", + "build_routed_experts", + "build_router", + "build_shared_expert", + "collect_moe_aux_loss", + "collect_moe_metrics", + "from_checkpoint_state_dict", + "get_score_function", + "is_moe_config", + "iter_moe_layers", + "plan_warm_start", + "register_balance_controller", + "register_experts_backend", + "register_moe_preset", + "register_router", + "register_score_function", + "register_shared_expert", + "register_state_dict_converter", + "resolve_moe_config", + "select_target_experts", + "to_checkpoint_state_dict", +] diff --git a/specforge/modeling/draft/moe/_registry.py b/specforge/modeling/draft/moe/_registry.py new file mode 100644 index 000000000..e1ef85e83 --- /dev/null +++ b/specforge/modeling/draft/moe/_registry.py @@ -0,0 +1,49 @@ +# coding=utf-8 +"""Tiny named-registry helper shared by the MoE component registries.""" + +from __future__ import annotations + +from typing import Callable, Dict, Generic, Optional, TypeVar + +T = TypeVar("T") + + +class Registry(Generic[T]): + """Name -> implementation map with a decorator form and helpful errors.""" + + def __init__(self, kind: str) -> None: + self.kind = kind + self._items: Dict[str, T] = {} + + def register(self, name: str, item: Optional[T] = None) -> Callable[[T], T] | T: + def _register(obj: T) -> T: + if name in self._items and self._items[name] is not obj: + raise ValueError(f"{self.kind} {name!r} is already registered") + self._items[name] = obj + return obj + + return _register if item is None else _register(item) + + def unregister(self, name: str) -> None: + self._items.pop(name, None) + + def get(self, name: str) -> T: + try: + return self._items[name] + except KeyError: + available = ", ".join(sorted(self._items)) or "" + raise KeyError( + f"unknown {self.kind} {name!r}; available: {available}" + ) from None + + def names(self) -> list[str]: + return sorted(self._items) + + def __contains__(self, name: object) -> bool: + return name in self._items + + def __getitem__(self, name: str) -> T: + return self.get(name) + + def __len__(self) -> int: + return len(self._items) diff --git a/specforge/modeling/draft/moe/balance.py b/specforge/modeling/draft/moe/balance.py new file mode 100644 index 000000000..3f45aa9d5 --- /dev/null +++ b/specforge/modeling/draft/moe/balance.py @@ -0,0 +1,79 @@ +# coding=utf-8 +"""Load balancing as a swappable policy. + +The controller sees routing outcomes and may (a) shift scores used for expert +*selection* (never the combine weights), (b) emit an auxiliary loss, and (c) +report load metrics. It is an ``nn.Module`` so implementations can own +buffers that travel with the checkpoint (e.g. a selection bias). + +Two timing rules every implementation must respect: + +- :meth:`observe` is called from the layer forward and must only *stash* + (overwrite, never accumulate): an activation-checkpoint recompute re-runs the + forward and must leave identical state behind. An auxiliary loss built here + is consumed by the trainer right after the same forward. +- :meth:`apply_pending_update` is called by the *model* before the next + forward, outside any checkpoint region, and may mutate selection state and + run collectives. Mutating selection state inside the forward would make the + recompute route differently and break checkpointing. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Dict, Optional, Type, Union + +import torch +from torch import nn + +from ._registry import Registry +from .config import MoEConfig + +if TYPE_CHECKING: # pragma: no cover - import cycle with router.py + from .router import RoutingResult + +MetricValue = Union[torch.Tensor, float] + + +class BalanceController(nn.Module): + """No balancing (registered as ``"none"``); the base for real policies.""" + + def __init__(self, cfg: MoEConfig, n_experts: int) -> None: + super().__init__() + self.cfg = cfg + self.n_experts = n_experts + + def adjust_selection_scores(self, scores: torch.Tensor) -> torch.Tensor: + """Scores used to pick experts; combine weights still use the raw ones.""" + return scores + + def observe(self, routing: "RoutingResult") -> None: + """Stash this forward's routing outcome (training only). + + ``routing.counts`` feeds load statistics; ``routing.scores`` (when the + router provides it) lets a policy build a differentiable auxiliary + loss to return from :meth:`aux_loss` for the same forward.""" + + def apply_pending_update(self) -> None: + """Consume the stash; called by the model outside checkpoint regions.""" + + def aux_loss(self) -> Optional[torch.Tensor]: + """Scaled auxiliary loss for the last forward, or ``None``.""" + return None + + def metrics(self) -> Dict[str, MetricValue]: + """Scalar diagnostics (rank-local; the trainer DP-averages them).""" + return {} + + +BALANCE_CONTROLLERS: Registry[Type[BalanceController]] = Registry( + "MoE balance controller" +) +BALANCE_CONTROLLERS.register("none", BalanceController) + + +def register_balance_controller(name: str): + return BALANCE_CONTROLLERS.register(name) + + +def build_balance_controller(cfg: MoEConfig, n_experts: int) -> BalanceController: + return BALANCE_CONTROLLERS.get(cfg.balance)(cfg, n_experts) diff --git a/specforge/modeling/draft/moe/config.py b/specforge/modeling/draft/moe/config.py new file mode 100644 index 000000000..ee6c58d87 --- /dev/null +++ b/specforge/modeling/draft/moe/config.py @@ -0,0 +1,238 @@ +# coding=utf-8 +"""MoE architecture config: the knobs a draft JSON can set, and their presets. + +Two kinds of keys, deliberately kept apart: + +- **Architecture keys** live at the top level of the draft JSON under the + target checkpoints' native HF names (``n_routed_experts``, + ``moe_intermediate_size``, ``scoring_func``, ...). They determine which + weights a checkpoint carries and how serving must route, so a draft JSON can + be assembled by copying them from the target's ``config.json``. A + ``moe_preset`` supplies the defaults for one target family; explicit keys + override the preset for ablations. +- **Training-only keys** live under ``dflash_config`` with an ``moe_`` prefix + (``moe_bias_update_rate``, ``moe_aux_loss_coeff``, ``moe_dispatch``). They + never change the checkpoint and are invisible to serving. +""" + +from __future__ import annotations + +from dataclasses import dataclass, fields +from typing import Any, Dict, Mapping, Optional + +from ._registry import Registry + +_MISSING = object() + +#: Draft-JSON top-level keys that describe the MoE architecture. +ARCHITECTURE_KEYS = ( + "n_routed_experts", + "num_experts_per_tok", + "moe_intermediate_size", + "n_shared_experts", + "shared_expert_intermediate_size", + "scoring_func", + "norm_topk_prob", + "routed_scaling_factor", + "n_group", + "topk_group", + "swiglu_limit", + "router", + "balance", + "shared_expert", + "shared_expert_gate", + "experts_backend", +) + +#: ``dflash_config`` keys (training-only) -> MoEConfig field. +TRAINING_KEYS = { + "moe_bias_update_rate": "bias_update_rate", + "moe_aux_loss_coeff": "aux_loss_coeff", + "moe_dispatch": "dispatch", + "moe_freeze_experts": "freeze_experts", +} + + +@dataclass(frozen=True) +class MoEConfig: + """Resolved MoE configuration for one draft model (all layers share it).""" + + preset: str + n_routed_experts: int + num_experts_per_tok: int + moe_intermediate_size: int + n_shared_experts: int = 1 + shared_expert_intermediate_size: Optional[int] = None + # Routing recipe. ``scoring_func`` names a registered score function; + # ``router``/``balance`` name registered component implementations. + scoring_func: str = "softmax" + norm_topk_prob: bool = True + routed_scaling_factor: float = 1.0 + n_group: int = 1 + topk_group: int = 1 + router: str = "topk" + balance: str = "none" + # Expert MLPs. + swiglu_limit: float = 0.0 + experts_backend: str = "grouped" + shared_expert: str = "swiglu" + shared_expert_gate: str = "none" + # Training-only (never part of the checkpoint). + bias_update_rate: float = 0.0 + aux_loss_coeff: float = 0.0 + dispatch: str = "sorted_loop" + #: Keep the routed experts fixed (e.g. warm-started from the target) and + #: train only the router, shared expert and the rest of the draft. Frozen + #: experts are replicated by the FSDP backend instead of sharded, which + #: removes the per-micro-batch weight all-gathers that dominate large MoEs. + freeze_experts: bool = False + + def __post_init__(self) -> None: + if self.n_routed_experts <= 0: + raise ValueError("n_routed_experts must be positive for an MoE FFN") + if not 0 < self.num_experts_per_tok <= self.n_routed_experts: + raise ValueError( + "num_experts_per_tok must be in [1, n_routed_experts], got " + f"{self.num_experts_per_tok} with {self.n_routed_experts} experts" + ) + if self.moe_intermediate_size <= 0: + raise ValueError("moe_intermediate_size must be positive") + if self.n_shared_experts not in (0, 1): + raise ValueError( + "n_shared_experts must be 0 or 1 (one shared expert of " + "shared_expert_intermediate_size width), got " + f"{self.n_shared_experts}" + ) + if self.shared_expert_intermediate_size is None: + object.__setattr__( + self, "shared_expert_intermediate_size", self.moe_intermediate_size + ) + if self.shared_expert_intermediate_size <= 0: + raise ValueError("shared_expert_intermediate_size must be positive") + if self.n_group <= 0 or self.n_routed_experts % self.n_group: + raise ValueError( + f"n_group={self.n_group} must divide n_routed_experts=" + f"{self.n_routed_experts}" + ) + if not 0 < self.topk_group <= self.n_group: + raise ValueError( + f"topk_group={self.topk_group} must be in [1, n_group={self.n_group}]" + ) + if self.swiglu_limit < 0: + raise ValueError("swiglu_limit must be >= 0 (0 disables the clamp)") + if self.bias_update_rate < 0 or self.aux_loss_coeff < 0: + raise ValueError("bias_update_rate and aux_loss_coeff must be >= 0") + + @property + def group_limited(self) -> bool: + """Whether routing restricts top-k to ``topk_group`` of ``n_group``.""" + return self.topk_group < self.n_group + + def as_dict(self) -> Dict[str, Any]: + return {f.name: getattr(self, f.name) for f in fields(self)} + + def serving_fields(self) -> Dict[str, Any]: + """The resolved recipe in the DeepSeek HF config vocabulary. + + Exports write these to ``config.json`` so a serving engine reads the + complete routing recipe without knowing SpecForge presets. Only the + keys a DeepSeek-style MoE reads; ``swiglu_limit`` is omitted when the + clamp is off (a serving engine treats 0 as a clamp at 0). + """ + out: Dict[str, Any] = { + "n_routed_experts": self.n_routed_experts, + "num_experts_per_tok": self.num_experts_per_tok, + "moe_intermediate_size": self.moe_intermediate_size, + "n_shared_experts": self.n_shared_experts, + "scoring_func": self.scoring_func, + "norm_topk_prob": self.norm_topk_prob, + "routed_scaling_factor": self.routed_scaling_factor, + "n_group": self.n_group, + "topk_group": self.topk_group, + "topk_method": "noaux_tc" if self.balance == "noaux_tc" else "greedy", + } + if self.swiglu_limit > 0: + out["swiglu_limit"] = self.swiglu_limit + return out + + +#: preset name -> architecture defaults (a subset of ARCHITECTURE_KEYS). +MOE_PRESETS: Registry[Dict[str, Any]] = Registry("MoE preset") + +_FIELD_NAMES = {f.name for f in fields(MoEConfig)} + + +def register_moe_preset(name: str, **defaults: Any) -> Dict[str, Any]: + """Register the architecture defaults of one target family. + + A preset may set any ``MoEConfig`` field except the training-only ones and + the per-run sizes (``n_routed_experts``, ``num_experts_per_tok``, + ``moe_intermediate_size``), which the draft JSON must state explicitly. + """ + forbidden = set(TRAINING_KEYS.values()) | { + "preset", + "n_routed_experts", + "num_experts_per_tok", + "moe_intermediate_size", + } + bad = sorted(set(defaults) - _FIELD_NAMES) + if bad: + raise ValueError(f"preset {name!r} sets unknown MoEConfig fields: {bad}") + bad = sorted(set(defaults) & forbidden) + if bad: + raise ValueError(f"preset {name!r} may not set per-run/training fields: {bad}") + MOE_PRESETS.register(name, dict(defaults)) + return defaults + + +def available_moe_presets() -> list[str]: + return MOE_PRESETS.names() + + +def _get(config: Any, key: str, default: Any = _MISSING) -> Any: + if isinstance(config, Mapping): + return config.get(key, default) + return getattr(config, key, default) + + +def is_moe_config(config: Any) -> bool: + """True when a draft config asks for an MoE FFN (``n_routed_experts > 0``).""" + value = _get(config, "n_routed_experts", 0) + return int(value or 0) > 0 + + +def resolve_moe_config(config: Any) -> Optional[MoEConfig]: + """Resolve a draft config (HF ``PretrainedConfig`` or dict) to ``MoEConfig``. + + Returns ``None`` for dense drafts. For MoE drafts, ``moe_preset`` is + required: it names the target family's routing recipe and is the only way + the architecture keys get validated defaults. + """ + if not is_moe_config(config): + return None + preset = _get(config, "moe_preset", None) + if not preset: + raise ValueError( + "n_routed_experts > 0 requires moe_preset in the draft config; " + f"available presets: {available_moe_presets() or ''}" + ) + values: Dict[str, Any] = dict(MOE_PRESETS.get(preset)) + for key in ARCHITECTURE_KEYS: + explicit = _get(config, key, _MISSING) + if explicit is not _MISSING and explicit is not None: + values[key] = explicit + dflash_config = _get(config, "dflash_config", None) or {} + unknown = sorted( + key + for key in dflash_config + if key.startswith("moe_") and key not in TRAINING_KEYS + ) + if unknown: + raise ValueError( + f"unknown MoE training keys in dflash_config: {unknown}; " + f"known: {sorted(TRAINING_KEYS)}" + ) + for json_key, field_name in TRAINING_KEYS.items(): + if json_key in dflash_config: + values[field_name] = dflash_config[json_key] + return MoEConfig(preset=preset, **values) diff --git a/specforge/modeling/draft/moe/deepseek_v4_target.py b/specforge/modeling/draft/moe/deepseek_v4_target.py new file mode 100644 index 000000000..17a76ed36 --- /dev/null +++ b/specforge/modeling/draft/moe/deepseek_v4_target.py @@ -0,0 +1,130 @@ +# coding=utf-8 +"""Read one DeepSeek-V4 target MoE layer as a warm-start source. + +The DeepSeek-V4-Flash checkpoints store routed experts as packed FP4 e2m1 +(two values per int8, low nibble first) with per-32 ue8m0 scales, the shared +expert and other linears as FP8 E4M3 with 128x128-block ue8m0 scales, the +gate in bf16 and the ``noaux_tc`` bias in fp32. This module dequantizes one +layer's ``ffn.*`` tensors to bf16 in the official naming +:func:`specforge.modeling.draft.moe.init.apply_warm_start` consumes +(``experts.{i}.w{1,2,3}.weight``, ``gate.weight``, ``gate.bias``, +``shared_experts.w{1,2,3}.weight``). Conventions follow the reference +``inference/convert.py``. + +Layers below ``num_hash_layers`` route by token hash (``gate.tid2eid``) and +carry no learned gate; they are rejected as warm-start sources. +""" + +from __future__ import annotations + +import json +import os +from typing import Dict, Iterable, Mapping, Tuple + +import torch + +FP4_TABLE = torch.tensor( + [ + 0.0, + 0.5, + 1.0, + 1.5, + 2.0, + 3.0, + 4.0, + 6.0, + 0.0, + -0.5, + -1.0, + -1.5, + -2.0, + -3.0, + -4.0, + -6.0, + ], + dtype=torch.float32, +) +FP8_BLOCK = 128 +FP4_GROUP = 32 + + +def dequant_fp8_block(weight: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: + """FP8 E4M3 ``[out, in]`` with e8m0 scale ``[out/128, in/128]`` -> bf16.""" + out_dim, in_dim = weight.shape + if out_dim % FP8_BLOCK or in_dim % FP8_BLOCK: + raise ValueError(f"fp8 weight {tuple(weight.shape)} is not 128-block aligned") + w = weight.float().view( + out_dim // FP8_BLOCK, FP8_BLOCK, in_dim // FP8_BLOCK, FP8_BLOCK + ) + w = w * scale.float()[:, None, :, None] + return w.view(out_dim, in_dim).to(torch.bfloat16) + + +def dequant_fp4_packed(weight: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: + """Packed FP4 int8 ``[out, in/2]`` with e8m0 scale ``[out, in/32]`` -> bf16. + + Low nibble = even element, high nibble = odd element.""" + if weight.dtype != torch.int8: + raise TypeError(f"packed fp4 weight must be int8, got {weight.dtype}") + out_dim, half_in = weight.shape + in_dim = half_in * 2 + x = weight.view(torch.uint8) + decoded = torch.stack( + [FP4_TABLE[(x & 0x0F).long()], FP4_TABLE[((x >> 4) & 0x0F).long()]], dim=-1 + ).view(out_dim, in_dim // FP4_GROUP, FP4_GROUP) + decoded = decoded * scale.float()[:, :, None] + return decoded.view(out_dim, in_dim).to(torch.bfloat16) + + +def dequantize_ffn_tensors( + raw: Mapping[str, torch.Tensor], prefix: str +) -> Dict[str, torch.Tensor]: + """``{prefix}...`` raw tensors of one MoE layer -> official-relative bf16 dict.""" + out: Dict[str, torch.Tensor] = {} + for name, tensor in raw.items(): + if not name.startswith(prefix) or name.endswith(".scale"): + continue + rel = name[len(prefix) :] + if tensor.dtype == torch.float8_e4m3fn: + value = dequant_fp8_block(tensor, raw[name[: -len(".weight")] + ".scale"]) + elif tensor.dtype == torch.int8: + value = dequant_fp4_packed(tensor, raw[name[: -len(".weight")] + ".scale"]) + elif rel == "gate.bias": + value = tensor.float() + elif rel == "gate.tid2eid": + raise ValueError(f"{prefix} is a hash-routed layer (no learned gate)") + else: + value = tensor.to(torch.bfloat16) + out[rel] = value + if "gate.bias" not in out or "gate.weight" not in out: + raise ValueError( + f"{prefix} has no learned gate; pick a layer >= num_hash_layers" + ) + return out + + +def _iter_layer_tensors( + snapshot_dir: str, prefix: str +) -> Iterable[Tuple[str, torch.Tensor]]: + from safetensors.torch import safe_open + + index = json.load(open(os.path.join(snapshot_dir, "model.safetensors.index.json"))) + weight_map = index["weight_map"] + shards = sorted({v for k, v in weight_map.items() if k.startswith(prefix)}) + if not shards: + raise KeyError(f"no tensors with prefix {prefix!r} in {snapshot_dir}") + for shard in shards: + with safe_open( + os.path.join(snapshot_dir, shard), framework="pt", device="cpu" + ) as h: + for name in h.keys(): + if name.startswith(prefix): + yield name, h.get_tensor(name) + + +def load_target_moe_layer(snapshot_dir: str, layer_id: int) -> Dict[str, torch.Tensor]: + """Dequantized ``layers.{layer_id}.ffn.*`` of a DeepSeek-V4 checkpoint dir.""" + prefix = f"layers.{layer_id}.ffn." + return dequantize_ffn_tensors( + dict(_iter_layer_tensors(snapshot_dir, prefix)), prefix + ) diff --git a/specforge/modeling/draft/moe/experts.py b/specforge/modeling/draft/moe/experts.py new file mode 100644 index 000000000..4d29cacee --- /dev/null +++ b/specforge/modeling/draft/moe/experts.py @@ -0,0 +1,57 @@ +# coding=utf-8 +"""Routed experts contract: the expert weights and how tokens reach them. + +An implementation owns the parameters of all ``n_routed_experts`` experts and +turns a :class:`RoutingResult` into the combined routed output. Weight layout +is the implementation's choice (per-expert modules, stacked ``[E, out, in]`` +tensors, ...); it must register a :mod:`.state_dict` converter if its native +layout differs from the official checkpoint naming +(``experts.{i}.w{1,2,3}.weight``). ``MoEConfig.dispatch`` is the +implementation's execution knob (e.g. sorted-segment loop vs grouped GEMM). +""" + +from __future__ import annotations + +import abc +from typing import Type + +import torch +from torch import nn + +from ._registry import Registry +from .config import MoEConfig +from .router import RoutingResult + + +class RoutedExperts(nn.Module, abc.ABC): + #: The training backend keeps a fully frozen instance replicated (outside + #: FSDP sharding): no weight all-gathers or gradient reduce-scatters. + fsdp_replicate_when_frozen = True + + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: + super().__init__() + self.cfg = cfg + self.hidden_size = hidden_size + self.n_experts = cfg.n_routed_experts + self.intermediate_size = cfg.moe_intermediate_size + + @abc.abstractmethod + def forward(self, x: torch.Tensor, routing: RoutingResult) -> torch.Tensor: + """``x`` is ``[T, hidden]``; return the combined routed output ``[T, hidden]`` + in ``x.dtype``.""" + + @abc.abstractmethod + def reset_parameters(self, std: float) -> None: + """Initialize expert weights with the draft's ``initializer_range`` so an + MoE FFN starts from the same distribution as the dense MLP it replaces.""" + + +EXPERTS_BACKENDS: Registry[Type[RoutedExperts]] = Registry("MoE experts backend") + + +def register_experts_backend(name: str): + return EXPERTS_BACKENDS.register(name) + + +def build_routed_experts(cfg: MoEConfig, hidden_size: int) -> RoutedExperts: + return EXPERTS_BACKENDS.get(cfg.experts_backend)(cfg, hidden_size) diff --git a/specforge/modeling/draft/moe/grouped_experts.py b/specforge/modeling/draft/moe/grouped_experts.py new file mode 100644 index 000000000..7a511a60e --- /dev/null +++ b/specforge/modeling/draft/moe/grouped_experts.py @@ -0,0 +1,154 @@ +# coding=utf-8 +"""Routed experts as three stacked parameters with sorted-segment dispatch. + +Weights live as ``w1``/``w2``/``w3`` of shape ``[E, out, in]``: grouped GEMMs +read them directly (a per-call ``torch.stack`` of hundreds of expert weights +would allocate a transient multi-GiB tensor) and FSDP ``use_orig_params`` +tracks 3 tensors instead of ``3*E``. Checkpoint FILES keep the official +per-expert naming (``experts.{i}.w{1,2,3}.weight``) through the converter +registered below. + +Dispatch (``MoEConfig.dispatch``): + +- ``"sorted_loop"``: one stable argsort turns routing into contiguous + per-expert segments, then one small GEMM per active expert. A per-expert + ``torch.where`` loop scales launch and autograd overhead with the number of + ACTIVE experts (~2x step time once the balancer spreads load). +- ``"grouped_mm"``: the same segments through ``torch._grouped_mm`` with + on-device offsets (no host sync). Used on CUDA when available; falls back to + the loop elsewhere. Same math up to bf16 rounding. +""" + +from __future__ import annotations + +import re + +import torch +import torch.nn.functional as F +from torch import nn + +from .config import MoEConfig +from .experts import RoutedExperts, register_experts_backend +from .router import RoutingResult +from .state_dict import register_state_dict_converter + +DISPATCH_MODES = ("sorted_loop", "grouped_mm") + + +def swiglu_clamped(gate: torch.Tensor, up: torch.Tensor, limit: float) -> torch.Tensor: + """SwiGLU in fp32 with the DeepSeek-V4 activation clamp (``limit`` 0 = off).""" + gate = gate.float() + up = up.float() + if limit > 0: + up = torch.clamp(up, min=-limit, max=limit) + gate = torch.clamp(gate, max=limit) + return F.silu(gate) * up + + +@register_experts_backend("grouped") +class GroupedExperts(RoutedExperts): + _WEIGHT_NAMES = ("w1", "w2", "w3") + + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: + super().__init__(cfg, hidden_size) + if cfg.dispatch not in DISPATCH_MODES: + raise ValueError( + f"unknown MoE dispatch {cfg.dispatch!r}; choose from {DISPATCH_MODES}" + ) + self.grouped_mm = cfg.dispatch == "grouped_mm" and hasattr(torch, "_grouped_mm") + self.swiglu_limit = float(cfg.swiglu_limit) + e, d, i = self.n_experts, hidden_size, self.intermediate_size + self.w1 = nn.Parameter(torch.empty(e, i, d)) + self.w2 = nn.Parameter(torch.empty(e, d, i)) + self.w3 = nn.Parameter(torch.empty(e, i, d)) + + def reset_parameters(self, std: float) -> None: + if self.w1.device.type == "meta": + return + for name in self._WEIGHT_NAMES: + nn.init.normal_(getattr(self, name), mean=0.0, std=std) + + def forward(self, x: torch.Tensor, routing: RoutingResult) -> torch.Tensor: + flat_expert = routing.indices.flatten() # [T*k] + order = flat_expert.argsort(stable=True) + token_of = order // routing.topk # routed token index per sorted slot + x_sorted = x.index_select(0, token_of) + w_sorted = routing.weights.reshape(-1, 1).index_select(0, order).float() + counts = routing.counts + + if self.grouped_mm and x.is_cuda: + offs = counts.cumsum(0).to(torch.int32) + gate = torch._grouped_mm(x_sorted, self.w1.transpose(-1, -2), offs=offs) + up = torch._grouped_mm(x_sorted, self.w3.transpose(-1, -2), offs=offs) + h = w_sorted * swiglu_clamped(gate, up, self.swiglu_limit) + y_routed = torch._grouped_mm( + h.to(x.dtype), self.w2.transpose(-1, -2), offs=offs + ) + else: + counts_list = counts.tolist() # one host sync per MoE forward + parts = [] + offset = 0 + for i, n in enumerate(counts_list): + if n == 0: + continue + seg = x_sorted[offset : offset + n] + h = w_sorted[offset : offset + n] * swiglu_clamped( + F.linear(seg, self.w1[i]), + F.linear(seg, self.w3[i]), + self.swiglu_limit, + ) + parts.append(F.linear(h.to(seg.dtype), self.w2[i])) + offset += n + if not parts: + return torch.zeros_like(x) + y_routed = torch.cat(parts, dim=0) + + y = torch.zeros(x.shape, dtype=torch.float32, device=x.device) + y = y.index_add(0, token_of, y_routed.float()) + return y.to(x.dtype) + + +_STACKED_KEY = re.compile(r"^(?P(?:.*\.)?experts)\.(?Pw[123])$") +_PER_EXPERT_KEY = re.compile( + r"^(?P(?:.*\.)?experts)\.(?P\d+)\.(?Pw[123])\.weight$" +) + + +def unstack_grouped_expert_state_dict(state: dict) -> dict: + """``experts.w1`` [E, out, in] -> ``experts.{i}.w1.weight``; no-op otherwise.""" + out = {} + for key, value in state.items(): + m = _STACKED_KEY.match(key) + if m is None or not isinstance(value, torch.Tensor) or value.dim() != 3: + out[key] = value + continue + for i in range(value.shape[0]): + out[f"{m['base']}.{i}.{m['w']}.weight"] = value[i] + return out + + +def stack_grouped_expert_state_dict(state: dict) -> dict: + """Inverse of :func:`unstack_grouped_expert_state_dict`.""" + groups: dict = {} + out = {} + for key, value in state.items(): + m = _PER_EXPERT_KEY.match(key) + if m is None: + out[key] = value + continue + groups.setdefault((m["base"], m["w"]), {})[int(m["idx"])] = value + for (base, w), members in groups.items(): + n = max(members) + 1 + if sorted(members) != list(range(n)): + raise KeyError( + f"{base}.*.{w}.weight is missing expert indices: have {sorted(members)}" + ) + out[f"{base}.{w}"] = torch.stack([members[i] for i in range(n)], dim=0) + return out + + +register_state_dict_converter( + "grouped_experts", + to_checkpoint=unstack_grouped_expert_state_dict, + from_checkpoint=stack_grouped_expert_state_dict, +) diff --git a/specforge/modeling/draft/moe/hooks.py b/specforge/modeling/draft/moe/hooks.py new file mode 100644 index 000000000..2867d370a --- /dev/null +++ b/specforge/modeling/draft/moe/hooks.py @@ -0,0 +1,63 @@ +# coding=utf-8 +"""Model-level plumbing for MoE layers. + +A draft model with MoE FFNs needs three things from its trainer loop, all +expressed here as functions over any ``nn.Module`` tree so DFlash, DFlash2 and +DSpark share them: + +- :func:`apply_pending_balance_updates` at the top of the model forward (in + training), outside activation-checkpoint regions; +- :func:`collect_moe_aux_loss` to add to the objective when a balance policy + emits one; +- :func:`collect_moe_metrics` for per-step diagnostics (``moe/...``). +""" + +from __future__ import annotations + +from typing import Dict, Iterator, Optional + +import torch +from torch import nn + +from .balance import MetricValue +from .layer import MoELayer + + +def iter_moe_layers(module: nn.Module) -> Iterator[MoELayer]: + for sub in module.modules(): + if isinstance(sub, MoELayer): + yield sub + + +def apply_pending_balance_updates(module: nn.Module) -> None: + for layer in iter_moe_layers(module): + layer.apply_pending_balance_update() + + +def collect_moe_aux_loss(module: nn.Module) -> Optional[torch.Tensor]: + """Sum of the layers' (already scaled) auxiliary losses, or ``None``.""" + total: Optional[torch.Tensor] = None + for layer in iter_moe_layers(module): + loss = layer.aux_loss() + if loss is None: + continue + total = loss if total is None else total + loss + return total + + +def collect_moe_metrics( + module: nn.Module, prefix: str = "moe/" +) -> Dict[str, MetricValue]: + """Layer-averaged scalar diagnostics; ``{}`` for dense models.""" + sums: Dict[str, MetricValue] = {} + n = 0 + for layer in iter_moe_layers(module): + n += 1 + for key, value in layer.metrics().items(): + sums[key] = value if key not in sums else sums[key] + value + if n == 0: + return {} + return { + f"{prefix}{key}": (value / n if isinstance(value, torch.Tensor) else value / n) + for key, value in sums.items() + } diff --git a/specforge/modeling/draft/moe/init.py b/specforge/modeling/draft/moe/init.py new file mode 100644 index 000000000..67f448280 --- /dev/null +++ b/specforge/modeling/draft/moe/init.py @@ -0,0 +1,106 @@ +# coding=utf-8 +"""Warm-start plans: which target experts seed which draft experts. + +A drafter whose MoE matches the target's expert shape can inherit expert +weights instead of training them from scratch. This module holds the +mapping (:func:`plan_warm_start`) and applying it to one ``MoELayer`` from a +target layer's *dequantized* tensors in official naming +(:func:`apply_warm_start`). Reading and dequantizing the target checkpoint is +target-specific and lives with the target's tooling. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import List, Mapping, Tuple + +import torch + +from .config import MoEConfig +from .layer import MoELayer +from .state_dict import from_checkpoint_state_dict + + +@dataclass(frozen=True) +class WarmStartPlan: + """``target_expert_ids[i]`` is the target expert that seeds draft expert ``i``.""" + + target_expert_ids: Tuple[int, ...] + copy_shared_expert: bool = True + copy_gate_rows: bool = True + + @property + def n_draft_experts(self) -> int: + return len(self.target_expert_ids) + + +def select_target_experts( + n_target: int, n_draft: int, strategy: str = "strided" +) -> Tuple[int, ...]: + """Pick ``n_draft`` distinct target experts. + + ``"strided"`` spreads picks evenly over the target's expert ids (the + default: no assumption about which experts matter for the draft's data); + ``"first"`` takes the leading ids. + """ + if n_draft <= 0 or n_target <= 0: + raise ValueError("expert counts must be positive") + if n_draft > n_target: + raise ValueError( + f"cannot seed {n_draft} draft experts from {n_target} target experts" + ) + if strategy == "first": + return tuple(range(n_draft)) + if strategy == "strided": + return tuple((i * n_target) // n_draft for i in range(n_draft)) + raise ValueError(f"unknown warm-start selection strategy {strategy!r}") + + +def plan_warm_start( + cfg: MoEConfig, n_target_experts: int, strategy: str = "strided" +) -> WarmStartPlan: + return WarmStartPlan( + target_expert_ids=select_target_experts( + n_target_experts, cfg.n_routed_experts, strategy + ), + copy_shared_expert=bool(cfg.n_shared_experts), + ) + + +_EXPERT_WEIGHTS = ("w1", "w2", "w3") + + +def apply_warm_start( + layer: MoELayer, plan: WarmStartPlan, source: Mapping[str, torch.Tensor] +) -> List[str]: + """Seed ``layer`` from one target MoE layer. + + ``source`` holds the target layer's tensors in official naming, relative + to the layer: ``experts.{j}.w{1,2,3}.weight``, ``gate.weight`` ``[E_t, H]``, + optionally ``gate.bias`` ``[E_t]`` and ``shared_experts.w{1,2,3}.weight``. + Returns the (module-native) keys that were loaded. + """ + if plan.n_draft_experts != layer.cfg.n_routed_experts: + raise ValueError( + f"plan seeds {plan.n_draft_experts} experts but the layer has " + f"{layer.cfg.n_routed_experts}" + ) + official = {} + for i, j in enumerate(plan.target_expert_ids): + for w in _EXPERT_WEIGHTS: + official[f"experts.{i}.{w}.weight"] = source[f"experts.{j}.{w}.weight"] + if plan.copy_gate_rows: + rows = torch.as_tensor(plan.target_expert_ids, dtype=torch.long) + official["gate.weight"] = source["gate.weight"][rows] + if "gate.bias" in source: + official["gate.bias"] = source["gate.bias"][rows] + if plan.copy_shared_expert and layer.shared_experts is not None: + for w in _EXPERT_WEIGHTS: + official[f"shared_experts.{w}.weight"] = source[ + f"shared_experts.{w}.weight" + ] + native = from_checkpoint_state_dict(official) + result = layer.load_state_dict(native, strict=False) + if result.unexpected_keys: + raise KeyError(f"warm start produced unexpected keys: {result.unexpected_keys}") + return sorted(native) diff --git a/specforge/modeling/draft/moe/layer.py b/specforge/modeling/draft/moe/layer.py new file mode 100644 index 000000000..9857d98dc --- /dev/null +++ b/specforge/modeling/draft/moe/layer.py @@ -0,0 +1,94 @@ +# coding=utf-8 +"""``MoELayer``: the FFN that composes router, experts and shared expert. + +Attribute names follow the official DeepSeek-style checkpoint layout +(``gate``, ``experts``, ``shared_experts``) so that per-implementation +converters only need to handle their own internals. +""" + +from __future__ import annotations + +from typing import Callable, Dict, Optional + +import torch +from torch import nn + +from .balance import MetricValue, build_balance_controller +from .config import MoEConfig, resolve_moe_config +from .experts import build_routed_experts +from .router import RoutingResult, build_router +from .shared import build_shared_expert + + +class MoELayer(nn.Module): + """Routed FFN: ``y = experts(x, gate(x)) + shared_experts(x)``.""" + + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: + super().__init__() + self.cfg = cfg + self.hidden_size = hidden_size + balance = build_balance_controller(cfg, cfg.n_routed_experts) + self.gate = build_router(cfg, hidden_size, balance) + self.experts = build_routed_experts(cfg, hidden_size) + if cfg.freeze_experts: + self.experts.requires_grad_(False) + self.shared_experts: Optional[nn.Module] = ( + build_shared_expert(cfg, hidden_size) if cfg.n_shared_experts else None + ) + # Detached per-expert counts of the last training forward, for metrics. + self.last_counts: Optional[torch.Tensor] = None + + @property + def balance(self): + return self.gate.balance + + def forward(self, x: torch.Tensor) -> torch.Tensor: + shape = x.shape + x = x.reshape(-1, self.hidden_size) + routing: RoutingResult = self.gate(x) + if self.training: + self.last_counts = routing.counts.detach() + self.balance.observe(routing) + y = self.experts(x, routing) + if self.shared_experts is not None: + y = y + self.shared_experts(x) + return y.view(shape) + + # -- model-level hooks (see hooks.py) --------------------------------- + def apply_pending_balance_update(self) -> None: + self.balance.apply_pending_update() + + def aux_loss(self) -> Optional[torch.Tensor]: + return self.balance.aux_loss() + + def metrics(self) -> Dict[str, MetricValue]: + out: Dict[str, MetricValue] = {} + counts = self.last_counts + if counts is not None and counts.numel(): + load = counts.float() + mean = load.mean().clamp_min(1e-9) + out["load_max_ratio"] = load.max() / mean + out["load_min_ratio"] = load.min() / mean + out["experts_unused_frac"] = (load == 0).float().mean() + out.update(self.balance.metrics()) + return out + + def reset_parameters(self, std: float) -> None: + """Initialize bare Parameters the HF ``_init_weights`` pass cannot see.""" + self.gate.reset_parameters(std) + self.experts.reset_parameters(std) + if self.shared_experts is not None: + self.shared_experts.reset_parameters(std) + + +def build_ffn(config, dense: Callable[[object], nn.Module]) -> nn.Module: + """The dense/MoE switch for a decoder layer's FFN. + + ``config`` is the draft's HF config; ``dense`` builds the dense MLP (the + kernel provider's factory) and is used verbatim when the config is dense, + so dense drafts are byte-for-byte unaffected by this package. + """ + moe_cfg = resolve_moe_config(config) + if moe_cfg is None: + return dense(config) + return MoELayer(moe_cfg, int(config.hidden_size)) diff --git a/specforge/modeling/draft/moe/noaux_tc.py b/specforge/modeling/draft/moe/noaux_tc.py new file mode 100644 index 000000000..063aec649 --- /dev/null +++ b/specforge/modeling/draft/moe/noaux_tc.py @@ -0,0 +1,130 @@ +# coding=utf-8 +"""Aux-loss-free balancing (DeepSeek-V3/V4 ``noaux_tc``). + +A per-expert fp32 bias shifts the scores used for *selection*; combine weights +still use the raw scores. A sign controller moves the bias against the +all-reduced expert load, so it is updated by the trainer loop, not gradients. + +Optionally (``dflash_config.moe_aux_loss_coeff`` > 0) the controller also +emits DeepSeek-V3's complementary sequence-wise balance loss, +``coeff * sum_e f_e * P_e`` with ``f_e`` the (E/(k*T))-scaled routed-token +fraction and ``P_e`` the mean normalized affinity, computed over the +micro-batch. The bias alone only *reorders* selection; from-scratch drafters +whose early router inputs are near-identical (mask-token embeddings) collapse +onto a few experts while the gate logits grow faster than the bias can +follow. The differentiable term bounds the logit gaps so input-driven routing +can emerge. + +Checkpoint naming: the bias is stored as ``.gate.bias`` (the DeepSeek +native key SGLang maps onto ``e_score_correction_bias``); the module keeps it +at ``gate.balance.bias``, converted at the state-dict boundary. +""" + +from __future__ import annotations + +import re +from typing import Dict, Optional + +import torch + +from .balance import BalanceController, MetricValue, register_balance_controller +from .config import MoEConfig +from .router import RoutingResult +from .state_dict import register_state_dict_converter + + +@register_balance_controller("noaux_tc") +class NoAuxTCController(BalanceController): + def __init__(self, cfg: MoEConfig, n_experts: int) -> None: + super().__init__(cfg, n_experts) + self.update_rate = float(cfg.bias_update_rate) + self.aux_loss_coeff = float(cfg.aux_loss_coeff) + self.register_buffer("bias", torch.zeros(n_experts, dtype=torch.float32)) + self._pending_counts: Optional[torch.Tensor] = None + self._aux_loss: Optional[torch.Tensor] = None + self.last_load: Optional[torch.Tensor] = None + + def _apply(self, fn, recurse=True): + module = super()._apply(fn, recurse) + # Sign-controller steps (~1e-3) vanish under bf16 rounding once the + # bias grows; keep the buffer fp32 through module-wide dtype casts. + if module.bias.dtype != torch.float32: + module.bias.data = module.bias.data.float() + return module + + def adjust_selection_scores(self, scores: torch.Tensor) -> torch.Tensor: + return scores + self.bias + + def observe(self, routing: RoutingResult) -> None: + # Overwrite, never accumulate: a checkpoint recompute re-runs the + # forward and must leave identical state behind. + self._pending_counts = routing.counts + self._aux_loss = None + scores = routing.scores + if self.aux_loss_coeff <= 0 or scores is None or not scores.requires_grad: + return + tokens, n_experts = scores.shape + # f_e: routed fraction scaled so that a uniform load gives 1. + f = routing.counts.float() * (n_experts / (routing.topk * tokens)) + # P_e: mean normalized affinity (the differentiable side). + p = (scores / scores.sum(dim=-1, keepdim=True).clamp_min(1e-20)).mean(0) + self._aux_loss = self.aux_loss_coeff * (f * p).sum() + + def aux_loss(self) -> Optional[torch.Tensor]: + return self._aux_loss + + def apply_pending_update(self) -> None: + import torch.distributed as dist + + counts = self._pending_counts + self._pending_counts = None + if counts is None or self.update_rate <= 0: + return + load = counts.float() + if dist.is_available() and dist.is_initialized(): + dist.all_reduce(load) + self.last_load = load + error = load.mean() - load + with torch.no_grad(): + self.bias += self.update_rate * torch.sign(error) + + def metrics(self) -> Dict[str, MetricValue]: + out: Dict[str, MetricValue] = {"bias_abs_max": self.bias.abs().max()} + if self._aux_loss is not None: + out["aux_loss"] = self._aux_loss.detach() + if self.last_load is not None: + mean = self.last_load.mean().clamp_min(1e-9) + out["global_load_max_ratio"] = self.last_load.max() / mean + out["global_load_min_ratio"] = self.last_load.min() / mean + return out + + +_NATIVE_BIAS = re.compile(r"^(?P(?:.*\.)?)gate\.balance\.bias$") +_OFFICIAL_BIAS = re.compile(r"^(?P(?:.*\.)?)gate\.bias$") + + +def _to_checkpoint(state: dict) -> dict: + return { + (f"{m['base']}gate.bias" if (m := _NATIVE_BIAS.match(k)) else k): v + for k, v in state.items() + } + + +def _is_moe_layer(state: dict, base: str) -> bool: + return f"{base}experts.w1" in state or f"{base}experts.0.w1.weight" in state + + +def _from_checkpoint(state: dict) -> dict: + out = {} + for key, value in state.items(): + m = _OFFICIAL_BIAS.match(key) + # Only an MoE layer's gate: a dense module named ``gate`` keeps its bias. + if m is not None and _is_moe_layer(state, m["base"]): + key = f"{m['base']}gate.balance.bias" + out[key] = value + return out + + +register_state_dict_converter( + "noaux_tc_bias", to_checkpoint=_to_checkpoint, from_checkpoint=_from_checkpoint +) diff --git a/specforge/modeling/draft/moe/presets.py b/specforge/modeling/draft/moe/presets.py new file mode 100644 index 000000000..3d0bc25d3 --- /dev/null +++ b/specforge/modeling/draft/moe/presets.py @@ -0,0 +1,27 @@ +# coding=utf-8 +"""Target-family MoE presets. + +A preset is the routing recipe of one target family; the draft JSON adds the +per-run sizes (``n_routed_experts``, ``num_experts_per_tok``, +``moe_intermediate_size``) and may override any key for ablations. +""" + +from .config import register_moe_preset + +# DeepSeek-V4 (e.g. DeepSeek-V4-Flash): sqrt(softplus) scores, aux-loss-free +# top-k with the sign-controlled selection bias, renormalized combine weights +# scaled by 1.5, one ungated shared expert, SwiGLU clamped at 10. The target's +# n_group == topk_group, so group-limited routing is off by default. +register_moe_preset( + "deepseek_v4", + scoring_func="sqrtsoftplus", + norm_topk_prob=True, + routed_scaling_factor=1.5, + n_shared_experts=1, + swiglu_limit=10.0, + router="topk", + balance="noaux_tc", + experts_backend="grouped", + shared_expert="swiglu", + shared_expert_gate="none", +) diff --git a/specforge/modeling/draft/moe/router.py b/specforge/modeling/draft/moe/router.py new file mode 100644 index 000000000..757bcc441 --- /dev/null +++ b/specforge/modeling/draft/moe/router.py @@ -0,0 +1,93 @@ +# coding=utf-8 +"""Router contract: hidden states -> per-token expert choices + combine weights. + +A router owns the gate projection and composes a :class:`BalanceController` +(``self.balance``) that may shift scores for *selection only*. Concrete routers +register by name (``MoEConfig.router``); score functions register separately +(``MoEConfig.scoring_func``) so one top-k router serves softmax, sigmoid and +sqrtsoftplus families. +""" + +from __future__ import annotations + +import abc +from dataclasses import dataclass +from typing import Callable, Optional, Type + +import torch +from torch import nn + +from ._registry import Registry +from .balance import BalanceController +from .config import MoEConfig + + +@dataclass +class RoutingResult: + """Routing decision for a flat batch of ``T`` tokens. + + ``weights`` are the final combine weights (normalized and scaled as the + recipe dictates) in fp32; ``indices`` the chosen experts; ``counts`` the + per-expert token counts on device (no host sync), which dispatch and + balancing both consume. ``scores`` are the full pre-selection affinities + (differentiable, fp32) for balance policies that need a gradient signal, + e.g. an auxiliary balance loss; routers may leave it ``None``. + """ + + weights: torch.Tensor # [T, k] fp32 + indices: torch.Tensor # [T, k] long + counts: torch.Tensor # [E] long + scores: Optional[torch.Tensor] = None # [T, E] fp32, differentiable + + @property + def topk(self) -> int: + return int(self.indices.shape[-1]) + + +#: name -> f(logits [T, E] fp32) -> scores [T, E] fp32 +SCORE_FUNCTIONS: Registry[Callable[[torch.Tensor], torch.Tensor]] = Registry( + "MoE score function" +) + + +def register_score_function(name: str): + return SCORE_FUNCTIONS.register(name) + + +def get_score_function(name: str) -> Callable[[torch.Tensor], torch.Tensor]: + return SCORE_FUNCTIONS.get(name) + + +class Router(nn.Module, abc.ABC): + """Base router. Subclasses implement :meth:`forward` and :meth:`reset_parameters`.""" + + def __init__( + self, cfg: MoEConfig, hidden_size: int, balance: BalanceController + ) -> None: + super().__init__() + self.cfg = cfg + self.hidden_size = hidden_size + self.n_experts = cfg.n_routed_experts + self.topk = cfg.num_experts_per_tok + self.balance = balance + + @abc.abstractmethod + def forward(self, x: torch.Tensor) -> RoutingResult: + """Route a flat ``[T, hidden]`` batch.""" + + @abc.abstractmethod + def reset_parameters(self, std: float) -> None: + """Initialize the gate weights (called from the model's ``_init_weights``).""" + + +ROUTERS: Registry[Type[Router]] = Registry("MoE router") + + +def register_router(name: str): + return ROUTERS.register(name) + + +def build_router( + cfg: MoEConfig, hidden_size: int, balance: BalanceController +) -> Router: + return ROUTERS.get(cfg.router)(cfg, hidden_size, balance) diff --git a/specforge/modeling/draft/moe/shared.py b/specforge/modeling/draft/moe/shared.py new file mode 100644 index 000000000..b655b577a --- /dev/null +++ b/specforge/modeling/draft/moe/shared.py @@ -0,0 +1,46 @@ +# coding=utf-8 +"""Shared expert contract: an always-on FFN added to the routed output. + +Families differ in the gate on it (DeepSeek: none; Qwen: a per-token sigmoid +gate), which is ``MoEConfig.shared_expert_gate``, and in its width +(``shared_expert_intermediate_size``). Implementations register by name +(``MoEConfig.shared_expert``). Checkpoint naming follows the official +``shared_experts.w{1,2,3}.weight`` layout unless a converter says otherwise. +""" + +from __future__ import annotations + +import abc +from typing import Type + +import torch +from torch import nn + +from ._registry import Registry +from .config import MoEConfig + + +class SharedExpert(nn.Module, abc.ABC): + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: + super().__init__() + self.cfg = cfg + self.hidden_size = hidden_size + self.intermediate_size = cfg.shared_expert_intermediate_size + + @abc.abstractmethod + def forward(self, x: torch.Tensor) -> torch.Tensor: + """``x`` is ``[T, hidden]``; return ``[T, hidden]`` in ``x.dtype``.""" + + def reset_parameters(self, std: float) -> None: + """Hook for bare Parameters; ``nn.Linear`` children are covered by HF init.""" + + +SHARED_EXPERTS: Registry[Type[SharedExpert]] = Registry("MoE shared expert") + + +def register_shared_expert(name: str): + return SHARED_EXPERTS.register(name) + + +def build_shared_expert(cfg: MoEConfig, hidden_size: int) -> SharedExpert: + return SHARED_EXPERTS.get(cfg.shared_expert)(cfg, hidden_size) diff --git a/specforge/modeling/draft/moe/state_dict.py b/specforge/modeling/draft/moe/state_dict.py new file mode 100644 index 000000000..020a43694 --- /dev/null +++ b/specforge/modeling/draft/moe/state_dict.py @@ -0,0 +1,56 @@ +# coding=utf-8 +"""Checkpoint-naming boundary for MoE modules. + +Checkpoint FILES (trainer checkpoints, warm-start sources, exports) use the +official per-expert naming so SGLang and HF loaders read them unchanged. +Modules may use a different native layout (e.g. stacked ``[E, out, in]`` +expert tensors, which FSDP ``use_orig_params`` and grouped GEMMs want). + +FSDP's full-state-dict hooks index the gathered dict by the module's own +parameter FQNs, so the rename cannot live inside ``state_dict()``; it lives at +the save/load boundary instead. Every place that reads a model's state for a +file, or loads a file into a model, goes through :func:`to_checkpoint_state_dict` +/ :func:`from_checkpoint_state_dict`: the training backend, warm start, and +the HF/SGLang exporters. Implementations with a native layout register a +converter pair; both directions must be no-ops on dicts already in the other +form, and on dense models. +""" + +from __future__ import annotations + +from typing import Callable, Dict, List, Tuple + +Converter = Callable[[Dict[str, object]], Dict[str, object]] + +_CONVERTERS: List[Tuple[str, Converter, Converter]] = [] + + +def register_state_dict_converter( + name: str, *, to_checkpoint: Converter, from_checkpoint: Converter +) -> None: + for existing, _, _ in _CONVERTERS: + if existing == name: + raise ValueError(f"state-dict converter {name!r} is already registered") + _CONVERTERS.append((name, to_checkpoint, from_checkpoint)) + + +def unregister_state_dict_converter(name: str) -> None: + _CONVERTERS[:] = [entry for entry in _CONVERTERS if entry[0] != name] + + +def registered_state_dict_converters() -> List[str]: + return [name for name, _, _ in _CONVERTERS] + + +def to_checkpoint_state_dict(state: Dict[str, object]) -> Dict[str, object]: + """Module-native naming -> official checkpoint naming (identity for dense).""" + for _, to_checkpoint, _ in _CONVERTERS: + state = to_checkpoint(state) + return state + + +def from_checkpoint_state_dict(state: Dict[str, object]) -> Dict[str, object]: + """Official checkpoint naming -> module-native naming (identity for dense).""" + for _, _, from_checkpoint in reversed(_CONVERTERS): + state = from_checkpoint(state) + return state diff --git a/specforge/modeling/draft/moe/swiglu_shared.py b/specforge/modeling/draft/moe/swiglu_shared.py new file mode 100644 index 000000000..e2a06061c --- /dev/null +++ b/specforge/modeling/draft/moe/swiglu_shared.py @@ -0,0 +1,30 @@ +# coding=utf-8 +"""Ungated SwiGLU shared expert (DeepSeek layout ``shared_experts.w{1,2,3}``).""" + +from __future__ import annotations + +import torch +from torch import nn + +from .config import MoEConfig +from .grouped_experts import swiglu_clamped +from .shared import SharedExpert, register_shared_expert + + +@register_shared_expert("swiglu") +class SwiGLUSharedExpert(SharedExpert): + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: + super().__init__(cfg, hidden_size) + if cfg.shared_expert_gate != "none": + raise ValueError( + "shared_expert='swiglu' is ungated; shared_expert_gate=" + f"{cfg.shared_expert_gate!r} needs a gated shared-expert implementation" + ) + self.swiglu_limit = float(cfg.swiglu_limit) + self.w1 = nn.Linear(hidden_size, self.intermediate_size, bias=False) + self.w2 = nn.Linear(self.intermediate_size, hidden_size, bias=False) + self.w3 = nn.Linear(hidden_size, self.intermediate_size, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + h = swiglu_clamped(self.w1(x), self.w3(x), self.swiglu_limit) + return self.w2(h.to(x.dtype)) diff --git a/specforge/modeling/draft/moe/topk_router.py b/specforge/modeling/draft/moe/topk_router.py new file mode 100644 index 000000000..c4376d632 --- /dev/null +++ b/specforge/modeling/draft/moe/topk_router.py @@ -0,0 +1,87 @@ +# coding=utf-8 +"""Top-k router with pluggable score functions and optional group-limited routing.""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F +from torch import nn + +from .balance import BalanceController +from .config import MoEConfig +from .router import ( + Router, + RoutingResult, + get_score_function, + register_router, + register_score_function, +) + + +@register_score_function("softmax") +def _softmax(logits: torch.Tensor) -> torch.Tensor: + return logits.softmax(dim=-1) + + +@register_score_function("sigmoid") +def _sigmoid(logits: torch.Tensor) -> torch.Tensor: + return torch.sigmoid(logits) + + +@register_score_function("sqrtsoftplus") +def _sqrtsoftplus(logits: torch.Tensor) -> torch.Tensor: + """DeepSeek-V4 scoring: ``sqrt(softplus(logits))``.""" + return F.softplus(logits).sqrt() + + +def group_limited_mask( + selection: torch.Tensor, n_group: int, topk_group: int +) -> torch.Tensor: + """Keep only the ``topk_group`` groups with the highest top-2 score sums + (DeepSeek ``noaux_tc`` group scoring); other groups become ``-inf``.""" + tokens, n_experts = selection.shape + grouped = selection.view(tokens, n_group, n_experts // n_group) + group_scores = grouped.topk(min(2, grouped.shape[-1]), dim=-1).values.sum(-1) + keep = group_scores.topk(topk_group, dim=-1).indices + mask = torch.zeros_like(group_scores, dtype=torch.bool).scatter_(1, keep, True) + return grouped.masked_fill(~mask.unsqueeze(-1), float("-inf")).view( + tokens, n_experts + ) + + +@register_router("topk") +class TopKRouter(Router): + """``scores = f(x W^T)``; pick top-k on balance-adjusted scores; combine + with the raw scores (optionally renormalized, then scaled).""" + + def __init__( + self, cfg: MoEConfig, hidden_size: int, balance: BalanceController + ) -> None: + super().__init__(cfg, hidden_size, balance) + self.weight = nn.Parameter(torch.empty(self.n_experts, hidden_size)) + self.score_fn = get_score_function(cfg.scoring_func) + + def reset_parameters(self, std: float) -> None: + nn.init.normal_(self.weight, mean=0.0, std=std) + + def forward(self, x: torch.Tensor) -> RoutingResult: + # Routing math in fp32 regardless of the model dtype. + scores = self.score_fn(F.linear(x.float(), self.weight.float())) + selection = self.balance.adjust_selection_scores(scores) + if self.cfg.group_limited: + selection = group_limited_mask( + selection, self.cfg.n_group, self.cfg.topk_group + ) + indices = selection.topk(self.topk, dim=-1).indices + weights = scores.gather(1, indices) + if self.cfg.norm_topk_prob: + weights = weights / (weights.sum(dim=-1, keepdim=True) + 1e-20) + weights = weights * self.cfg.routed_scaling_factor + flat = indices.flatten() + # scatter_add instead of bincount: CUDA bincount hides a device sync. + counts = torch.zeros( + self.n_experts, dtype=torch.long, device=x.device + ).scatter_add_(0, flat, torch.ones_like(flat)) + return RoutingResult( + weights=weights, indices=indices, counts=counts, scores=scores + ) diff --git a/specforge/training/backend.py b/specforge/training/backend.py index f4d673bb4..c488883aa 100644 --- a/specforge/training/backend.py +++ b/specforge/training/backend.py @@ -209,11 +209,21 @@ def _frozen_target_modules(model: nn.Module) -> tuple[nn.Module, ...]: all-gather them before every optimizer window without saving optimizer memory, which is the wrong trade-off for the current trainer recipes. """ + candidates = [ + module + for name in ("lm_head", "embed_tokens") + if isinstance(module := getattr(model, name, None), nn.Module) + ] + # Submodules that opt in (e.g. frozen MoE experts warm-started from the + # target) are replicated as well: sharding tens of GB of frozen weights + # would re-gather them on every micro-batch for no optimizer savings. + candidates += [ + module + for module in model.modules() + if getattr(module, "fsdp_replicate_when_frozen", False) + ] modules = [] - for name in ("lm_head", "embed_tokens"): - module = getattr(model, name, None) - if not isinstance(module, nn.Module): - continue + for module in candidates: parameters = tuple(module.parameters()) if parameters and not any( parameter.requires_grad for parameter in parameters @@ -411,22 +421,30 @@ def _full_state_ctx(self, state_dict_config=None): ) def _module_state_dict(self) -> dict: + # Checkpoint FILES use the official parameter naming; modules may use + # a different native layout (MoE experts). Convert at this boundary: + # FSDP's full-state-dict hooks need the module's own FQNs. + from specforge.modeling.draft.moe import to_checkpoint_state_dict + if self._wrapper_kind == "ddp": if dist.is_initialized() and dist.get_rank() != 0: return {} - return self.module.module.state_dict() + return to_checkpoint_state_dict(self.module.module.state_dict()) if self._wrapper_kind != "fsdp": - return self.module.state_dict() + return to_checkpoint_state_dict(self.module.state_dict()) from torch.distributed.fsdp import FullStateDictConfig # gather to rank0 CPU only — materializing the full model on every # rank's GPU is wasted memory when only rank0 writes it. cfg = FullStateDictConfig(offload_to_cpu=True, rank0_only=True) with self._full_state_ctx(cfg): - return self.module.state_dict() + return to_checkpoint_state_dict(self.module.state_dict()) def _load_module_state_dict(self, model_state: dict) -> None: + from specforge.modeling.draft.moe import from_checkpoint_state_dict + # every rank loads the full state dict read from the shared file. + model_state = from_checkpoint_state_dict(model_state) if self._wrapper_kind == "ddp": self.module.module.load_state_dict(model_state) return diff --git a/specforge/training/model_loading.py b/specforge/training/model_loading.py index a4f8d8616..38286ac83 100644 --- a/specforge/training/model_loading.py +++ b/specforge/training/model_loading.py @@ -379,7 +379,15 @@ def _load_pretrained_draft_state( cache_dir: Optional[str], trust_remote_code: bool, ) -> Dict[str, Any]: - from specforge.modeling.auto import AutoDraftModel + from specforge.modeling.auto import AutoDraftModel, load_native_state_dict + + # Drafts with a native module layout (MoE experts): read and regroup the + # files directly. Building a full second model per rank doubled the CPU + # footprint (~65 GB each for a 256-expert drafter) and OOM-killed 8-rank + # trainers during warm start. + native = load_native_state_dict(draft_config, source) + if native is not None: + return {key: value.detach().cpu() for key, value in native.items()} loaded, loading_info = AutoDraftModel.from_pretrained( source, @@ -398,6 +406,25 @@ def _load_pretrained_draft_state( return state +def _rank_staggered(fn): + """Run ``fn`` on one distributed rank at a time (all ranks call this).""" + try: + import torch.distributed as dist + + active = dist.is_available() and dist.is_initialized() + except ImportError: # pragma: no cover + active = False + if not active or dist.get_world_size() == 1: + return fn() + rank, world = dist.get_rank(), dist.get_world_size() + result = None + for turn in range(world): + if turn == rank: + result = fn() + dist.barrier() + return result + + def warm_start_draft_model( model: Any, source: str, @@ -411,27 +438,43 @@ def warm_start_draft_model( """Load only draft weights, never optimizer/counters/RNG training state.""" runtime_state = _runtime_state_file(source) - if runtime_state is not None: - checkpoint_format: Literal["specforge", "pretrained"] = "specforge" - state = _load_specforge_draft_state(runtime_state, expected_strategy=strategy) - else: - checkpoint_format = "pretrained" - state = _load_pretrained_draft_state( - source, - draft_config=draft_config, - cache_dir=cache_dir, - trust_remote_code=trust_remote_code, - ) + checkpoint_format: Literal["specforge", "pretrained"] = ( + "specforge" if runtime_state is not None else "pretrained" + ) - if not state: - raise ValueError(f"warm-start checkpoint {source!r} contains no draft weights") - try: - result = model.load_state_dict(state, strict=False) - except RuntimeError as exc: - raise ValueError( - f"warm-start checkpoint {source!r} has incompatible draft tensor " - f"shapes: {exc}" - ) from exc + def _load() -> Tuple[Dict[str, Any], Any]: + if runtime_state is not None: + state = _load_specforge_draft_state( + runtime_state, expected_strategy=strategy + ) + else: + state = _load_pretrained_draft_state( + source, + draft_config=draft_config, + cache_dir=cache_dir, + trust_remote_code=trust_remote_code, + ) + if not state: + raise ValueError( + f"warm-start checkpoint {source!r} contains no draft weights" + ) + try: + # Files use the official naming; modules may use a native MoE layout. + from specforge.modeling.draft.moe import from_checkpoint_state_dict + + result = model.load_state_dict( + from_checkpoint_state_dict(state), strict=False + ) + except RuntimeError as exc: + raise ValueError( + f"warm-start checkpoint {source!r} has incompatible draft tensor " + f"shapes: {exc}" + ) from exc + return state, result + + # Every rank loads the full draft state; for large drafts the transient + # host memory of N simultaneous loads can exceed the node, so take turns. + state, result = _rank_staggered(_load) loaded_keys = len(state) - len(result.unexpected_keys) if result.unexpected_keys or loaded_keys == 0: raise ValueError( diff --git a/specforge/training/strategies/base.py b/specforge/training/strategies/base.py index f592ccdc0..0f1ddf526 100644 --- a/specforge/training/strategies/base.py +++ b/specforge/training/strategies/base.py @@ -56,6 +56,27 @@ class StepContext: collect_detailed_metrics: bool = True +def _moe_metrics(model_wrapper: nn.Module) -> Dict[str, Any]: + """``moe/...`` load diagnostics of the wrapped draft; ``{}`` for dense drafts.""" + from specforge.modeling.draft.moe import collect_moe_metrics + + draft_model = getattr(model_wrapper, "draft_model", None) + if draft_model is None: + return {} + return collect_moe_metrics(draft_model) + + +def _with_moe_aux_loss(model_wrapper: nn.Module, loss: torch.Tensor) -> torch.Tensor: + """Add the MoE balance policies' auxiliary loss (if any) for this forward.""" + from specforge.modeling.draft.moe import collect_moe_aux_loss + + draft_model = getattr(model_wrapper, "draft_model", None) + if draft_model is None: + return loss + aux = collect_moe_aux_loss(draft_model) + return loss if aux is None else loss + aux.to(loss.dtype) + + def linear_lambda_base( global_step: int, total_steps: int, @@ -546,8 +567,9 @@ def forward_loss( metrics["accuracy_denom"] = model_metrics["accuracy_denom"] if "selector_loss_alpha" in model_metrics: metrics["selector_loss_alpha"] = model_metrics["selector_loss_alpha"] + metrics.update(_moe_metrics(self.dflash_model)) return StepOutput( - loss=loss, + loss=_with_moe_aux_loss(self.dflash_model, loss), metrics=metrics, ratio_metrics=model_metrics.get("ratio_metrics", {}), loss_terms=model_metrics.get("loss_terms"), @@ -616,8 +638,9 @@ def forward_loss( ): if name in model_metrics: metrics[name] = model_metrics[name] + metrics.update(_moe_metrics(self.dspark_model)) return StepOutput( - loss=loss, + loss=_with_moe_aux_loss(self.dspark_model, loss), metrics=metrics, ratio_metrics=ratio_metrics, ) diff --git a/tests/test_modeling/test_moe.py b/tests/test_modeling/test_moe.py new file mode 100644 index 000000000..e8745c7bb --- /dev/null +++ b/tests/test_modeling/test_moe.py @@ -0,0 +1,468 @@ +# coding=utf-8 +"""MoE FFN skeleton: config resolution, component contracts, composition, and +the model/trainer/checkpoint seams — exercised with stub components so the +contracts are pinned independently of any target-family implementation.""" + +import unittest +from types import SimpleNamespace + +import torch +from torch import nn +from transformers import Qwen3Config +from transformers.models.qwen3.modeling_qwen3 import Qwen3MLP + +from specforge.modeling.draft.dflash import DFlashDraftModel +from specforge.modeling.draft.moe import ( + BALANCE_CONTROLLERS, + EXPERTS_BACKENDS, + MOE_PRESETS, + ROUTERS, + SCORE_FUNCTIONS, + SHARED_EXPERTS, + BalanceController, + MoEConfig, + MoELayer, + RoutedExperts, + Router, + RoutingResult, + SharedExpert, + apply_pending_balance_updates, + build_ffn, + collect_moe_aux_loss, + collect_moe_metrics, + from_checkpoint_state_dict, + iter_moe_layers, + plan_warm_start, + register_balance_controller, + register_experts_backend, + register_moe_preset, + register_router, + register_score_function, + register_shared_expert, + register_state_dict_converter, + resolve_moe_config, + select_target_experts, + to_checkpoint_state_dict, +) +from specforge.modeling.draft.moe.state_dict import unregister_state_dict_converter +from specforge.training.strategies.base import _moe_metrics + +PRESET = "_test_family" + + +# --- stub components: a minimal but real top-k reference for the contracts --- + + +class _StubBalance(BalanceController): + def __init__(self, cfg, n_experts): + super().__init__(cfg, n_experts) + self.register_buffer("bias", torch.zeros(n_experts)) + self.observed = [] + self.applied = 0 + + def adjust_selection_scores(self, scores): + return scores + self.bias + + def observe(self, routing): + self.observed.append(routing.counts.detach().clone()) + self.saw_scores = routing.scores is not None + + def apply_pending_update(self): + self.applied += 1 + + def aux_loss(self): + if self.cfg.aux_loss_coeff <= 0: + return None + return torch.tensor(self.cfg.aux_loss_coeff) + + def metrics(self): + return {"stub_metric": 1.0} + + +class _StubRouter(Router): + def __init__(self, cfg, hidden_size, balance): + super().__init__(cfg, hidden_size, balance) + self.weight = nn.Parameter(torch.empty(self.n_experts, hidden_size)) + self.score_fn = SCORE_FUNCTIONS.get(cfg.scoring_func) + + def reset_parameters(self, std): + nn.init.normal_(self.weight, std=std) + + def forward(self, x): + scores = self.score_fn(x.float() @ self.weight.float().t()) + indices = self.balance.adjust_selection_scores(scores).topk(self.topk).indices + weights = scores.gather(1, indices) + if self.cfg.norm_topk_prob: + weights = weights / weights.sum(-1, keepdim=True) + weights = weights * self.cfg.routed_scaling_factor + counts = torch.zeros( + self.n_experts, dtype=torch.long, device=x.device + ).scatter_add_(0, indices.flatten(), torch.ones_like(indices.flatten())) + return RoutingResult( + weights=weights, indices=indices, counts=counts, scores=scores + ) + + +class _StubExperts(RoutedExperts): + def __init__(self, cfg, hidden_size): + super().__init__(cfg, hidden_size) + self.w = nn.Parameter(torch.empty(self.n_experts, hidden_size, hidden_size)) + + def reset_parameters(self, std): + nn.init.normal_(self.w, std=std) + + def forward(self, x, routing): + # dense gather: fine for tiny tests, pins the [T, k] -> [T, H] contract + per_choice = torch.einsum( + "td,tkdo->tko", x.float(), self.w[routing.indices].float() + ) + return (routing.weights.unsqueeze(-1) * per_choice).sum(1).to(x.dtype) + + +class _StubShared(SharedExpert): + def __init__(self, cfg, hidden_size): + super().__init__(cfg, hidden_size) + self.proj = nn.Linear(hidden_size, hidden_size, bias=False) + + def forward(self, x): + return self.proj(x) + + +def setUpModule(): + register_score_function("_test_softplus")(torch.nn.functional.softplus) + register_router("_test_router")(_StubRouter) + register_balance_controller("_test_balance")(_StubBalance) + register_experts_backend("_test_experts")(_StubExperts) + register_shared_expert("_test_shared")(_StubShared) + register_moe_preset( + PRESET, + scoring_func="_test_softplus", + router="_test_router", + balance="_test_balance", + experts_backend="_test_experts", + shared_expert="_test_shared", + routed_scaling_factor=1.5, + ) + + +def tearDownModule(): + SCORE_FUNCTIONS.unregister("_test_softplus") + ROUTERS.unregister("_test_router") + BALANCE_CONTROLLERS.unregister("_test_balance") + EXPERTS_BACKENDS.unregister("_test_experts") + SHARED_EXPERTS.unregister("_test_shared") + MOE_PRESETS.unregister(PRESET) + + +def _moe_json(**overrides): + payload = dict( + hidden_size=16, + moe_preset=PRESET, + n_routed_experts=8, + num_experts_per_tok=2, + moe_intermediate_size=8, + dflash_config={}, + ) + payload.update(overrides) + return payload + + +def _dflash_config(moe=True, **overrides): + fields = dict( + architectures=["DFlashDraftModel"], + block_size=2, + hidden_size=16, + intermediate_size=32, + num_attention_heads=2, + num_key_value_heads=1, + num_hidden_layers=2, + num_target_layers=6, + head_dim=8, + max_position_embeddings=64, + vocab_size=32, + layer_types=["full_attention", "full_attention"], + dflash_config={"attention_mode": "gqa"}, + ) + if moe: + fields.update(_moe_json(hidden_size=16)) + fields["dflash_config"] = { + "attention_mode": "gqa", + "moe_bias_update_rate": 0.005, + } + fields.update(overrides) + config = Qwen3Config(**fields) + config._attn_implementation = "sdpa" + return config + + +class TestMoEConfig(unittest.TestCase): + def test_dense_config_resolves_to_none(self): + self.assertIsNone(resolve_moe_config({"hidden_size": 16})) + self.assertIsNone(resolve_moe_config(SimpleNamespace(n_routed_experts=0))) + + def test_preset_is_required_for_moe(self): + with self.assertRaisesRegex(ValueError, "moe_preset"): + resolve_moe_config(_moe_json(moe_preset=None)) + with self.assertRaisesRegex(KeyError, "available"): + resolve_moe_config(_moe_json(moe_preset="no-such-family")) + + def test_preset_defaults_and_explicit_overrides(self): + cfg = resolve_moe_config(_moe_json()) + self.assertEqual(cfg.preset, PRESET) + self.assertEqual(cfg.scoring_func, "_test_softplus") + self.assertEqual(cfg.routed_scaling_factor, 1.5) + self.assertEqual( + cfg.shared_expert_intermediate_size, 8 + ) # defaults to moe width + cfg = resolve_moe_config( + _moe_json(routed_scaling_factor=2.0, shared_expert_intermediate_size=4) + ) + self.assertEqual(cfg.routed_scaling_factor, 2.0) + self.assertEqual(cfg.shared_expert_intermediate_size, 4) + # Works on attribute-style configs (HF PretrainedConfig) as well. + self.assertEqual(resolve_moe_config(_dflash_config()).n_routed_experts, 8) + + def test_training_knobs_come_from_dflash_config(self): + cfg = resolve_moe_config( + _moe_json( + dflash_config={ + "moe_bias_update_rate": 0.01, + "moe_dispatch": "grouped_mm", + } + ) + ) + self.assertEqual(cfg.bias_update_rate, 0.01) + self.assertEqual(cfg.dispatch, "grouped_mm") + with self.assertRaisesRegex(ValueError, "unknown MoE training keys"): + resolve_moe_config(_moe_json(dflash_config={"moe_bias_udpate_rate": 1})) + + def test_validation(self): + with self.assertRaisesRegex(ValueError, "num_experts_per_tok"): + resolve_moe_config(_moe_json(num_experts_per_tok=9)) + with self.assertRaisesRegex(ValueError, "n_shared_experts"): + resolve_moe_config(_moe_json(n_shared_experts=2)) + with self.assertRaisesRegex(ValueError, "n_group"): + resolve_moe_config(_moe_json(n_group=3)) + cfg = resolve_moe_config(_moe_json(n_group=4, topk_group=2)) + self.assertTrue(cfg.group_limited) + self.assertFalse(resolve_moe_config(_moe_json()).group_limited) + + def test_presets_cannot_set_per_run_or_training_fields(self): + with self.assertRaisesRegex(ValueError, "per-run/training"): + register_moe_preset("_bad", n_routed_experts=4) + with self.assertRaisesRegex(ValueError, "unknown MoEConfig fields"): + register_moe_preset("_bad", nonsense=1) + self.assertNotIn("_bad", MOE_PRESETS) + + def test_registry_errors_name_the_kind_and_choices(self): + with self.assertRaisesRegex(KeyError, "MoE router.*available"): + ROUTERS.get("missing") + with self.assertRaisesRegex(ValueError, "already registered"): + register_router("_test_router")(_StubRouter.__mro__[1]) + self.assertIn("none", BALANCE_CONTROLLERS) + + +class TestMoELayerComposition(unittest.TestCase): + def _layer(self, **overrides): + torch.manual_seed(0) + cfg = resolve_moe_config(_moe_json(**overrides)) + layer = MoELayer(cfg, 16) + layer.reset_parameters(std=0.02) + return layer + + def test_dense_config_uses_the_dense_factory_verbatim(self): + sentinel = nn.Identity() + self.assertIs( + build_ffn(SimpleNamespace(hidden_size=16), lambda c: sentinel), sentinel + ) + + def test_moe_config_builds_official_attribute_layout(self): + layer = build_ffn(SimpleNamespace(**_moe_json()), lambda c: nn.Identity()) + self.assertIsInstance(layer, MoELayer) + names = {name for name, _ in layer.named_children()} + self.assertEqual(names, {"gate", "experts", "shared_experts"}) + self.assertIsInstance(layer.gate.balance, _StubBalance) + self.assertIs(layer.balance, layer.gate.balance) + + def test_forward_preserves_shape_and_adds_shared_expert(self): + layer = self._layer() + x = torch.randn(2, 3, 16) + y = layer(x) + self.assertEqual(y.shape, x.shape) + routed_only = self._layer(n_shared_experts=0) + self.assertIsNone(routed_only.shared_experts) + routed_only.load_state_dict( + {k: v for k, v in layer.state_dict().items() if "shared" not in k} + ) + shared = layer.shared_experts(x.reshape(-1, 16)).view_as(x) + self.assertTrue(torch.allclose(y, routed_only(x) + shared, atol=1e-5)) + + def test_training_observes_counts_and_eval_does_not(self): + layer = self._layer().train() + x = torch.randn(5, 16) + layer(x) + self.assertEqual(len(layer.balance.observed), 1) + self.assertEqual(int(layer.balance.observed[0].sum()), 5 * 2) + self.assertTrue(torch.equal(layer.last_counts, layer.balance.observed[0])) + self.assertTrue(layer.balance.saw_scores) + layer.eval() + layer(x) + self.assertEqual(len(layer.balance.observed), 1) + + def test_model_hooks_delegate_to_the_controller(self): + layer = self._layer(dflash_config={"moe_aux_loss_coeff": 0.25}).train() + layer(torch.randn(4, 16)) + layer.apply_pending_balance_update() + self.assertEqual(layer.balance.applied, 1) + self.assertAlmostEqual(float(layer.aux_loss()), 0.25) + metrics = layer.metrics() + self.assertEqual( + set(metrics), + {"load_max_ratio", "load_min_ratio", "experts_unused_frac", "stub_metric"}, + ) + self.assertGreaterEqual(float(metrics["load_max_ratio"]), 1.0) + self.assertLessEqual(float(metrics["load_min_ratio"]), 1.0) + + def test_selection_bias_changes_choice_not_weights(self): + layer = self._layer().eval() + x = torch.randn(1, 16) + before = layer.gate(x) + layer.balance.bias[:] = -1e3 + favored = int((before.indices[0, 0] + 1) % 8) + layer.balance.bias[favored] = 0.0 + after = layer.gate(x) + self.assertIn(favored, after.indices[0].tolist()) + # combine weights come from raw scores: still normalized and scaled + self.assertAlmostEqual(float(after.weights.detach().sum()), 1.5, places=5) + + +class TestHooks(unittest.TestCase): + def _tree(self, **overrides): + cfg = resolve_moe_config(_moe_json(**overrides)) + a, b = MoELayer(cfg, 16), MoELayer(cfg, 16) + for layer in (a, b): + layer.reset_parameters(std=0.02) + return nn.Sequential(nn.Linear(16, 16), a, nn.Linear(16, 16), b), (a, b) + + def test_iteration_updates_and_aggregation(self): + tree, (a, b) = self._tree(dflash_config={"moe_aux_loss_coeff": 0.5}) + self.assertEqual(list(iter_moe_layers(tree)), [a, b]) + tree.train()(torch.randn(3, 16)) + apply_pending_balance_updates(tree) + self.assertEqual((a.balance.applied, b.balance.applied), (1, 1)) + self.assertAlmostEqual(float(collect_moe_aux_loss(tree)), 1.0) + metrics = collect_moe_metrics(tree) + self.assertTrue(all(key.startswith("moe/") for key in metrics)) + self.assertAlmostEqual(float(metrics["moe/stub_metric"]), 1.0) + + def test_dense_trees_are_inert(self): + dense = nn.Sequential(nn.Linear(16, 16)) + apply_pending_balance_updates(dense) + self.assertIsNone(collect_moe_aux_loss(dense)) + self.assertEqual(collect_moe_metrics(dense), {}) + self.assertEqual(_moe_metrics(SimpleNamespace()), {}) + + def test_strategy_metrics_read_the_wrapped_draft(self): + tree, _ = self._tree() + tree.train()(torch.randn(3, 16)) + metrics = _moe_metrics(SimpleNamespace(draft_model=tree)) + self.assertIn("moe/load_max_ratio", metrics) + + +class TestStateDictBoundary(unittest.TestCase): + def test_dense_state_passes_through_unchanged(self): + state = {"a": torch.zeros(1), "layers.0.mlp.gate_proj.weight": torch.ones(1)} + for convert in (to_checkpoint_state_dict, from_checkpoint_state_dict): + out = convert(dict(state)) + self.assertEqual(set(out), set(state)) + for key in state: + self.assertIs(out[key], state[key]) + + def test_registered_converters_apply_in_both_directions(self): + def to_ckpt(state): + return {k.replace("native.", "official."): v for k, v in state.items()} + + def from_ckpt(state): + return {k.replace("official.", "native."): v for k, v in state.items()} + + register_state_dict_converter( + "_test", to_checkpoint=to_ckpt, from_checkpoint=from_ckpt + ) + try: + with self.assertRaisesRegex(ValueError, "already registered"): + register_state_dict_converter( + "_test", to_checkpoint=to_ckpt, from_checkpoint=from_ckpt + ) + official = to_checkpoint_state_dict({"native.w": 1, "other": 2}) + self.assertEqual(official, {"official.w": 1, "other": 2}) + self.assertEqual( + from_checkpoint_state_dict(official), {"native.w": 1, "other": 2} + ) + finally: + unregister_state_dict_converter("_test") + self.assertEqual(to_checkpoint_state_dict({"native.w": 1}), {"native.w": 1}) + + +class TestWarmStartPlan(unittest.TestCase): + def test_selection_strategies(self): + self.assertEqual(select_target_experts(256, 4), (0, 64, 128, 192)) + self.assertEqual(select_target_experts(8, 3, "strided"), (0, 2, 5)) + self.assertEqual(select_target_experts(8, 3, "first"), (0, 1, 2)) + self.assertEqual(len(set(select_target_experts(256, 64))), 64) + with self.assertRaises(ValueError): + select_target_experts(4, 8) + with self.assertRaises(ValueError): + select_target_experts(8, 2, "random") + + def test_plan_follows_the_moe_config(self): + cfg = resolve_moe_config(_moe_json(n_shared_experts=0)) + plan = plan_warm_start(cfg, n_target_experts=64) + self.assertEqual(plan.n_draft_experts, 8) + self.assertFalse(plan.copy_shared_expert) + self.assertTrue(plan.copy_gate_rows) + + +class TestDFlashWiring(unittest.TestCase): + def _forward(self, model): + return model( + position_ids=torch.arange(6).unsqueeze(0), + noise_embedding=torch.randn(1, 2, 16), + target_hidden=torch.randn(1, 4, 2 * 16), + ) + + def test_dense_draft_is_unchanged(self): + model = DFlashDraftModel(_dflash_config(moe=False)) + for layer in model.layers: + self.assertIsInstance(layer.mlp, Qwen3MLP) + self.assertEqual(list(iter_moe_layers(model)), []) + + def test_moe_draft_layers_init_and_apply_balance_updates(self): + model = DFlashDraftModel(_dflash_config()) + layers = list(iter_moe_layers(model)) + self.assertEqual(len(layers), 2) + for layer in layers: + self.assertIs(layer, [m for m in model.layers if m.mlp is layer][0].mlp) + self.assertEqual(layer.cfg.bias_update_rate, 0.005) + # _init_weights reached the bare Parameters (no uninitialized memory) + self.assertTrue(torch.isfinite(layer.gate.weight).all()) + self.assertGreater(float(layer.gate.weight.abs().sum()), 0.0) + self.assertTrue(torch.isfinite(layer.experts.w).all()) + model.train() + self._forward(model) + self._forward(model) + self.assertEqual([layer.balance.applied for layer in layers], [2, 2]) + self.assertEqual([len(layer.balance.observed) for layer in layers], [2, 2]) + model.eval() + self._forward(model) + self.assertEqual([layer.balance.applied for layer in layers], [2, 2]) + + def test_state_dict_names_follow_the_official_layout(self): + model = DFlashDraftModel(_dflash_config()) + keys = set(model.state_dict()) + self.assertIn("layers.0.mlp.gate.weight", keys) + self.assertIn("layers.0.mlp.shared_experts.proj.weight", keys) + self.assertTrue(any(k.startswith("layers.0.mlp.experts.") for k in keys)) + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/tests/test_modeling/test_moe_deepseek_v4.py b/tests/test_modeling/test_moe_deepseek_v4.py new file mode 100644 index 000000000..e26f28c4c --- /dev/null +++ b/tests/test_modeling/test_moe_deepseek_v4.py @@ -0,0 +1,599 @@ +# coding=utf-8 +"""DeepSeek-V4 MoE preset: routing, aux-loss-free balancing, grouped experts, +official checkpoint naming, warm start, and DFlash integration.""" + +import json +import unittest +from pathlib import Path + +import torch +from torch import nn +from transformers import Qwen3Config + +from specforge.modeling.draft.dflash import DFlashDraftModel +from specforge.modeling.draft.moe import ( + MOE_PRESETS, + MoELayer, + apply_pending_balance_updates, + apply_warm_start, + collect_moe_metrics, + from_checkpoint_state_dict, + get_score_function, + iter_moe_layers, + plan_warm_start, + resolve_moe_config, + to_checkpoint_state_dict, +) +from specforge.modeling.draft.moe.grouped_experts import ( + GroupedExperts, + stack_grouped_expert_state_dict, + swiglu_clamped, + unstack_grouped_expert_state_dict, +) +from specforge.modeling.draft.moe.noaux_tc import NoAuxTCController +from specforge.modeling.draft.moe.swiglu_shared import SwiGLUSharedExpert +from specforge.modeling.draft.moe.topk_router import TopKRouter, group_limited_mask + +REPO_ROOT = Path(__file__).resolve().parents[2] +CUDA = torch.cuda.is_available() + + +def _json(**overrides): + payload = dict( + moe_preset="deepseek_v4", + n_routed_experts=8, + num_experts_per_tok=2, + moe_intermediate_size=16, + dflash_config={"moe_bias_update_rate": 1e-3}, + ) + payload.update(overrides) + return payload + + +def _layer(**overrides) -> MoELayer: + torch.manual_seed(0) + layer = MoELayer(resolve_moe_config(_json(**overrides)), 32) + layer.reset_parameters(std=0.05) + for p in layer.shared_experts.parameters(): + nn.init.normal_(p, std=0.05) + return layer + + +def _dflash_config(**overrides): + fields = dict( + architectures=["DFlashDraftModel"], + block_size=2, + hidden_size=32, + intermediate_size=64, + num_attention_heads=2, + num_key_value_heads=1, + num_hidden_layers=2, + num_target_layers=6, + head_dim=16, + max_position_embeddings=64, + vocab_size=32, + layer_types=["full_attention", "full_attention"], + initializer_range=0.02, + **_json(dflash_config={"attention_mode": "gqa", "moe_bias_update_rate": 0.005}), + ) + fields.update(overrides) + config = Qwen3Config(**fields) + config._attn_implementation = "sdpa" + return config + + +def _reference_forward(layer: MoELayer, x: torch.Tensor) -> torch.Tensor: + """Per-token dense reference of the routed + shared FFN.""" + routing = layer.gate(x) + e = layer.experts + out = torch.zeros_like(x, dtype=torch.float32) + for t in range(x.shape[0]): + for k in range(routing.topk): + i = int(routing.indices[t, k]) + h = swiglu_clamped(x[t] @ e.w1[i].t(), x[t] @ e.w3[i].t(), e.swiglu_limit) + out[t] += routing.weights[t, k] * (h.to(x.dtype) @ e.w2[i].t()).float() + return (out + layer.shared_experts(x).float()).to(x.dtype) + + +class TestPresetAndConfig(unittest.TestCase): + def test_preset_matches_deepseek_v4_recipe(self): + self.assertIn("deepseek_v4", MOE_PRESETS) + cfg = resolve_moe_config(_json()) + self.assertEqual(cfg.scoring_func, "sqrtsoftplus") + self.assertEqual(cfg.balance, "noaux_tc") + self.assertEqual(cfg.routed_scaling_factor, 1.5) + self.assertTrue(cfg.norm_topk_prob) + self.assertEqual(cfg.swiglu_limit, 10.0) + self.assertEqual(cfg.n_shared_experts, 1) + self.assertFalse(cfg.group_limited) + + def test_checked_in_draft_config_resolves(self): + payload = json.loads( + (REPO_ROOT / "configs" / "deepseek-v4-flash-dspark-moe.json").read_text() + ) + cfg = resolve_moe_config(payload) + self.assertEqual( + (cfg.n_routed_experts, cfg.num_experts_per_tok, cfg.moe_intermediate_size), + (64, 6, 2048), + ) + self.assertEqual(cfg.dispatch, "grouped_mm") + self.assertEqual(cfg.bias_update_rate, 1e-3) + dense = json.loads( + (REPO_ROOT / "configs" / "deepseek-v4-flash-dspark.json").read_text() + ) + moe_only = { + "moe_preset", + "n_routed_experts", + "num_experts_per_tok", + "moe_intermediate_size", + "n_shared_experts", + } + self.assertEqual(set(payload) - set(dense), moe_only) + for key in dense: + if key != "dflash_config": + self.assertEqual(payload[key], dense[key], key) + + def test_score_functions(self): + logits = torch.tensor([[0.0, 2.0, -3.0]]) + self.assertTrue( + torch.allclose( + get_score_function("sqrtsoftplus")(logits), + torch.nn.functional.softplus(logits).sqrt(), + ) + ) + self.assertAlmostEqual( + float(get_score_function("softmax")(logits).sum()), 1.0, places=5 + ) + self.assertTrue( + torch.allclose(get_score_function("sigmoid")(logits), logits.sigmoid()) + ) + + +class TestTopKRouter(unittest.TestCase): + def test_weights_are_renormalized_and_scaled(self): + layer = _layer() + routing = layer.gate(torch.randn(5, 32)) + self.assertIsInstance(layer.gate, TopKRouter) + self.assertTrue(torch.allclose(routing.weights.sum(-1), torch.full((5,), 1.5))) + self.assertEqual(int(routing.counts.sum()), 10) + for row in routing.indices.tolist(): + self.assertEqual(len(set(row)), 2) + + def test_group_limited_routing_stays_within_selected_groups(self): + selection = torch.randn(64, 16) + masked = group_limited_mask(selection, n_group=4, topk_group=1) + finite = torch.isfinite(masked).view(64, 4, 4) + self.assertTrue((finite.all(-1).sum(-1) == 1).all()) + layer = _layer( + n_routed_experts=16, n_group=4, topk_group=1, num_experts_per_tok=3 + ) + routing = layer.gate(torch.randn(10, 32)) + groups = routing.indices // 4 + self.assertTrue((groups == groups[:, :1]).all()) + + +class TestNoAuxTC(unittest.TestCase): + def test_deferred_bias_update_semantics(self): + layer = _layer().train() + ctrl = layer.balance + self.assertIsInstance(ctrl, NoAuxTCController) + layer(torch.randn(6, 32)) + self.assertIsNotNone(ctrl._pending_counts) + before = ctrl.bias.clone() + layer.apply_pending_balance_update() + self.assertIsNone(ctrl._pending_counts) + self.assertFalse(torch.equal(before, ctrl.bias)) + self.assertTrue(((ctrl.bias - before).abs() <= 1e-3 + 1e-7).all()) + after = ctrl.bias.clone() + layer.apply_pending_balance_update() # nothing pending: no-op + self.assertTrue(torch.equal(after, ctrl.bias)) + layer.eval() + layer(torch.randn(6, 32)) + self.assertIsNone(ctrl._pending_counts) + + def test_bias_stays_fp32_and_moves_selection_only(self): + layer = _layer().to(torch.bfloat16).eval() + self.assertEqual(layer.balance.bias.dtype, torch.float32) + x = torch.randn(4, 32, dtype=torch.bfloat16) + self.assertEqual(layer(x).dtype, torch.bfloat16) + layer.balance.bias[:] = -100.0 + layer.balance.bias[3] = 0.0 + routing = layer.gate(x) + self.assertTrue((routing.indices == 3).any(dim=-1).all()) + self.assertTrue(torch.allclose(routing.weights.sum(-1), torch.full((4,), 1.5))) + + def test_aux_balance_loss_is_differentiable_and_uniform_at_balance(self): + layer = _layer(dflash_config={"moe_aux_loss_coeff": 0.5}).train() + x = torch.randn(64, 32) + layer(x) + aux = layer.aux_loss() + self.assertIsNotNone(aux) + self.assertTrue(aux.requires_grad) + aux.backward() + self.assertIsNotNone(layer.gate.weight.grad) + self.assertGreater(float(layer.gate.weight.grad.abs().sum()), 0.0) + # perfectly uniform routing and affinities give exactly coeff * 1 + routing = layer.gate(x) + n_experts = layer.cfg.n_routed_experts + uniform_scores = torch.full((64, n_experts), 0.25, requires_grad=True) + counts = torch.full((n_experts,), 64 * routing.topk // n_experts) + layer.balance.observe( + type(routing)(routing.weights, routing.indices, counts, uniform_scores) + ) + self.assertAlmostEqual(float(layer.balance.aux_loss()), 0.5, places=5) + self.assertIn("aux_loss", layer.balance.metrics()) + # disabled by default, and never built without a gradient signal + layer = _layer().train() + layer(torch.randn(8, 32)) + self.assertIsNone(layer.aux_loss()) + with torch.no_grad(): + aux_layer = _layer(dflash_config={"moe_aux_loss_coeff": 0.5}).train() + aux_layer(torch.randn(8, 32)) + self.assertIsNone(aux_layer.aux_loss()) + + def test_metrics_include_bias_and_global_load(self): + layer = _layer().train() + layer(torch.randn(6, 32)) + layer.apply_pending_balance_update() + metrics = collect_moe_metrics(nn.Sequential(layer)) + for key in ( + "moe/load_max_ratio", + "moe/bias_abs_max", + "moe/global_load_max_ratio", + ): + self.assertIn(key, metrics) + + +class TestGroupedExperts(unittest.TestCase): + def test_layout_init_and_dense_reference(self): + layer = _layer().eval() + e = layer.experts + self.assertIsInstance(e, GroupedExperts) + self.assertIsInstance(layer.shared_experts, SwiGLUSharedExpert) + self.assertEqual(tuple(e.w1.shape), (8, 16, 32)) + self.assertEqual(tuple(e.w2.shape), (8, 32, 16)) + self.assertAlmostEqual(float(e.w1.std()), 0.05, delta=0.01) + x = torch.randn(7, 32) + self.assertTrue( + torch.allclose(layer(x), _reference_forward(layer, x), atol=1e-5) + ) + + def test_swiglu_clamp(self): + gate = torch.tensor([50.0, -50.0]) + up = torch.tensor([50.0, -50.0]) + clamped = swiglu_clamped(gate, up, 10.0) + # gate clamps to max 10, up to [-10, 10]: silu(10)*10 and silu(-50)*-10 + expected = torch.nn.functional.silu(torch.tensor([10.0, -50.0])) * torch.tensor( + [10.0, -10.0] + ) + self.assertTrue(torch.allclose(clamped, expected, atol=1e-6)) + self.assertGreater(float(swiglu_clamped(gate, up, 0.0)[0]), 100.0) + + def test_unknown_dispatch_is_rejected(self): + with self.assertRaisesRegex(ValueError, "dispatch"): + _layer(dflash_config={"moe_dispatch": "magic"}) + + @unittest.skipUnless( + CUDA and hasattr(torch, "_grouped_mm"), "needs CUDA grouped GEMM" + ) + def test_grouped_mm_matches_sorted_loop(self): + torch.manual_seed(3) + layer = _layer(n_routed_experts=16).to("cuda", torch.bfloat16).train() + x = (torch.randn(6, 32, device="cuda") * 0.5).to(torch.bfloat16) + results = {} + for grouped in (False, True): + layer.experts.grouped_mm = grouped + layer.zero_grad(set_to_none=True) + xg = x.clone().requires_grad_(True) + y = layer(xg) + y.float().square().sum().backward() + results[grouped] = ( + y.detach().clone(), + xg.grad.clone(), + { + n: p.grad.clone() + for n, p in layer.named_parameters() + if p.grad is not None + }, + ) + (y0, dx0, g0), (y1, dx1, g1) = results[False], results[True] + self.assertTrue(torch.allclose(y0.float(), y1.float(), rtol=2e-2, atol=2e-2)) + self.assertTrue(torch.allclose(dx0.float(), dx1.float(), rtol=2e-2, atol=2e-2)) + self.assertLessEqual(set(g0), set(g1)) + for name in g1: + if name in g0: + self.assertTrue( + torch.allclose( + g0[name].float(), g1[name].float(), rtol=2e-2, atol=2e-2 + ), + name, + ) + else: + self.assertEqual(int(torch.count_nonzero(g1[name])), 0, name) + + +class TestCheckpointNaming(unittest.TestCase): + def test_layer_roundtrip_through_official_naming(self): + layer = _layer() + native = layer.state_dict() + self.assertIn("experts.w1", native) + self.assertIn("gate.balance.bias", native) + official = to_checkpoint_state_dict(native) + self.assertIn("experts.0.w1.weight", official) + self.assertIn("gate.bias", official) + self.assertIn("shared_experts.w1.weight", official) + self.assertFalse( + any("balance" in k or k.endswith("experts.w1") for k in official) + ) + fresh = _layer(dflash_config={"moe_bias_update_rate": 0.0}) + fresh.load_state_dict(from_checkpoint_state_dict(official), strict=True) + self.assertTrue(torch.equal(fresh.experts.w2, layer.experts.w2)) + # both directions are idempotent + self.assertEqual(set(to_checkpoint_state_dict(official)), set(official)) + self.assertEqual(set(from_checkpoint_state_dict(native)), set(native)) + + def test_dense_gate_bias_is_left_alone(self): + state = { + "head.gate.bias": torch.zeros(1), + "head.gate.weight": torch.zeros(1, 1), + } + self.assertEqual(set(from_checkpoint_state_dict(state)), set(state)) + + def test_stack_rejects_missing_expert_indices(self): + official = unstack_grouped_expert_state_dict(_layer().state_dict()) + del official["experts.3.w2.weight"] + with self.assertRaises(KeyError): + stack_grouped_expert_state_dict(official) + + +class TestWarmStart(unittest.TestCase): + def test_apply_plan_copies_selected_experts_gate_rows_and_shared(self): + layer = _layer() + n_target = 16 + source = { + "gate.weight": torch.randn(n_target, 32), + "gate.bias": torch.randn(n_target), + } + for j in range(n_target): + source[f"experts.{j}.w1.weight"] = torch.randn(16, 32) + source[f"experts.{j}.w2.weight"] = torch.randn(32, 16) + source[f"experts.{j}.w3.weight"] = torch.randn(16, 32) + for w, shape in (("w1", (16, 32)), ("w2", (32, 16)), ("w3", (16, 32))): + source[f"shared_experts.{w}.weight"] = torch.randn(*shape) + plan = plan_warm_start(layer.cfg, n_target_experts=n_target) + self.assertEqual(plan.target_expert_ids, (0, 2, 4, 6, 8, 10, 12, 14)) + loaded = apply_warm_start(layer, plan, source) + self.assertIn("experts.w1", loaded) + for i, j in enumerate(plan.target_expert_ids): + self.assertTrue( + torch.equal(layer.experts.w1[i], source[f"experts.{j}.w1.weight"]) + ) + self.assertTrue(torch.equal(layer.gate.weight[i], source["gate.weight"][j])) + self.assertEqual( + float(layer.balance.bias[i]), float(source["gate.bias"][j]) + ) + self.assertTrue( + torch.equal( + layer.shared_experts.w2.weight, source["shared_experts.w2.weight"] + ) + ) + with self.assertRaises(ValueError): + apply_warm_start(_layer(n_routed_experts=4), plan, source) + + +class TestServingExport(unittest.TestCase): + def test_serving_fields_carry_the_resolved_recipe(self): + fields = resolve_moe_config(_json()).serving_fields() + self.assertEqual(fields["topk_method"], "noaux_tc") + self.assertEqual(fields["scoring_func"], "sqrtsoftplus") + self.assertEqual(fields["routed_scaling_factor"], 1.5) + self.assertEqual(fields["swiglu_limit"], 10.0) + self.assertEqual((fields["n_group"], fields["topk_group"]), (1, 1)) + # a disabled clamp is omitted rather than exported as a clamp at 0 + self.assertNotIn( + "swiglu_limit", resolve_moe_config(_json(swiglu_limit=0)).serving_fields() + ) + + def test_hf_export_reloads_and_carries_serving_config(self): + import os + import tempfile + + from specforge.export import export_to_hf + from specforge.modeling.auto import AutoDraftModel + + torch.manual_seed(1) + config = _dflash_config() + model = DFlashDraftModel(config).to(torch.bfloat16) + for layer in iter_moe_layers(model): + layer.balance.bias.uniform_(-1.0, 1.0) + workdir = tempfile.mkdtemp(prefix="moe_export_") + config_path = os.path.join(workdir, "draft.json") + config.save_pretrained(workdir) + os.replace(os.path.join(workdir, "config.json"), config_path) + ckpt_dir = os.path.join(workdir, "run-step1") + os.makedirs(ckpt_dir) + torch.save( + { + "draft_state_dict": to_checkpoint_state_dict(model.state_dict()), + "strategy": "dflash", + "global_step": 1, + }, + os.path.join(ckpt_dir, "training_state.pt"), + ) + out = export_to_hf(ckpt_dir, config_path, os.path.join(workdir, "hf")) + exported = json.loads((Path(out) / "config.json").read_text()) + self.assertEqual(exported["topk_method"], "noaux_tc") + self.assertEqual(exported["scoring_func"], "sqrtsoftplus") + self.assertEqual(exported["n_routed_experts"], 8) + from safetensors import safe_open + + with safe_open(os.path.join(out, "model.safetensors"), "pt") as f: + keys = set(f.keys()) + self.assertIn("layers.0.mlp.experts.0.w1.weight", keys) + self.assertIn("layers.0.mlp.gate.bias", keys) + # HF from_pretrained assigns tensors by key; SpecForge's loader must + # convert the official naming back into the stacked module layout. + reloaded = AutoDraftModel.from_pretrained(out, torch_dtype=torch.bfloat16) + fresh = reloaded.state_dict() + for key, value in model.state_dict().items(): + self.assertTrue(torch.equal(value.float(), fresh[key].float()), key) + # the trainer's weights-only warm start reads the same directory + from specforge.training.model_loading import warm_start_draft_model + + target = DFlashDraftModel(_dflash_config()).to(torch.bfloat16) + report = warm_start_draft_model( + target, out, draft_config=config, strategy="dflash" + ) + self.assertEqual(report.checkpoint_format, "pretrained") + for key, value in model.state_dict().items(): + self.assertTrue( + torch.equal(value.float(), target.state_dict()[key].float()), key + ) + + +class TestFrozenExperts(unittest.TestCase): + def test_freeze_experts_trains_router_and_shared_only(self): + layer = _layer(dflash_config={"moe_freeze_experts": True}).train() + self.assertFalse(any(p.requires_grad for p in layer.experts.parameters())) + self.assertTrue(layer.gate.weight.requires_grad) + self.assertTrue(all(p.requires_grad for p in layer.shared_experts.parameters())) + y = layer(torch.randn(6, 32, requires_grad=True)) + y.float().sum().backward() + self.assertIsNotNone(layer.gate.weight.grad) + self.assertIsNone(layer.experts.w1.grad) + + def test_backend_replicates_frozen_experts(self): + from specforge.training.backend import FSDPTrainingBackend + + config = _dflash_config() + config.dflash_config = {**config.dflash_config, "moe_freeze_experts": True} + model = DFlashDraftModel(config) + ignored = FSDPTrainingBackend._frozen_target_modules(model) + self.assertEqual([type(m).__name__ for m in ignored], ["GroupedExperts"] * 2) + trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) + frozen = sum(p.numel() for m in ignored for p in m.parameters()) + self.assertGreater(frozen, trainable) + # trained experts stay sharded + self.assertEqual( + FSDPTrainingBackend._frozen_target_modules( + DFlashDraftModel(_dflash_config()) + ), + (), + ) + + +class TestDeepseekV4TargetDequant(unittest.TestCase): + def test_fp4_and_fp8_dequant_conventions(self): + from specforge.modeling.draft.moe.deepseek_v4_target import ( + FP4_TABLE, + dequant_fp4_packed, + dequant_fp8_block, + dequantize_ffn_tensors, + ) + + # one row of 32 fp4 values: nibbles 0..15 twice; low nibble = even index + codes = torch.arange(16, dtype=torch.uint8) + packed = ( + (codes | (codes << 4)).repeat(2).view(1, 32).to(torch.int8) + ) # 64 values + scale = torch.tensor([[2.0, 0.5]], dtype=torch.float32) # two groups of 32 + out = dequant_fp4_packed(packed, scale) + self.assertEqual(tuple(out.shape), (1, 64)) + expected = FP4_TABLE[codes.long()].repeat_interleave(2).repeat(2) + expected[:32] *= 2.0 + expected[32:] *= 0.5 + self.assertTrue(torch.equal(out.float()[0], expected)) + w8 = torch.full((128, 256), 1.0).to(torch.float8_e4m3fn) + s8 = torch.tensor([[1.0, 4.0]]) + d8 = dequant_fp8_block(w8, s8).float() + self.assertTrue((d8[:, :128] == 1.0).all() and (d8[:, 128:] == 4.0).all()) + raw = { + "layers.3.ffn.gate.weight": torch.randn(4, 8), + "layers.3.ffn.gate.bias": torch.randn(4), + "layers.3.ffn.experts.0.w1.weight": packed.repeat(2, 1), + "layers.3.ffn.experts.0.w1.scale": scale.repeat(2, 1), + "layers.3.ffn.shared_experts.w1.weight": w8, + "layers.3.ffn.shared_experts.w1.scale": s8, + } + rel = dequantize_ffn_tensors(raw, "layers.3.ffn.") + self.assertEqual( + set(rel), + { + "gate.weight", + "gate.bias", + "experts.0.w1.weight", + "shared_experts.w1.weight", + }, + ) + self.assertEqual(rel["gate.bias"].dtype, torch.float32) + self.assertEqual(rel["experts.0.w1.weight"].dtype, torch.bfloat16) + with self.assertRaisesRegex(ValueError, "hash-routed"): + dequantize_ffn_tensors( + { + "layers.1.ffn.gate.tid2eid": torch.zeros(4), + "layers.1.ffn.gate.weight": torch.zeros(4, 8), + }, + "layers.1.ffn.", + ) + + +class TestDFlashIntegration(unittest.TestCase): + def _forward(self, model): + return model( + position_ids=torch.arange(6).unsqueeze(0), + noise_embedding=torch.randn(1, 2, 32), + target_hidden=torch.randn(1, 4, 2 * 32), + ) + + def test_layers_train_and_balance_through_the_model(self): + model = DFlashDraftModel(_dflash_config()) + layers = list(iter_moe_layers(model)) + self.assertEqual(len(layers), 2) + for layer in layers: + self.assertIsInstance(layer.experts, GroupedExperts) + self.assertAlmostEqual(float(layer.experts.w1.std()), 0.02, delta=0.005) + self.assertAlmostEqual(float(layer.gate.weight.std()), 0.02, delta=0.005) + self.assertTrue(torch.equal(layer.balance.bias, torch.zeros(8))) + model.train() + out = self._forward(model) + out.float().square().mean().backward() + for layer in layers: + self.assertIsNotNone(layer.experts.w2.grad) + self.assertIsNotNone(layer.gate.weight.grad) + self.assertTrue(torch.equal(layer.balance.bias, torch.zeros(8))) # deferred + self._forward(model) # applies the pending update before routing + self.assertTrue(any(layer.balance.bias.abs().sum() > 0 for layer in layers)) + + def test_model_checkpoint_uses_official_naming_and_reloads(self): + model = DFlashDraftModel(_dflash_config()) + official = to_checkpoint_state_dict(model.state_dict()) + self.assertIn("layers.0.mlp.experts.0.w1.weight", official) + self.assertIn("layers.0.mlp.gate.bias", official) + self.assertIn("layers.1.mlp.shared_experts.w3.weight", official) + self.assertFalse( + any(".balance." in k or k.endswith(".experts.w1") for k in official) + ) + fresh = DFlashDraftModel(_dflash_config()) + fresh.load_state_dict(from_checkpoint_state_dict(official), strict=True) + self.assertTrue( + torch.equal(fresh.layers[1].mlp.experts.w1, model.layers[1].mlp.experts.w1) + ) + + def test_dense_config_is_unaffected(self): + config = _dflash_config() + for key in ( + "moe_preset", + "n_routed_experts", + "num_experts_per_tok", + "moe_intermediate_size", + "n_shared_experts", + ): + if hasattr(config, key): + delattr(config, key) + model = DFlashDraftModel(config) + self.assertEqual(list(iter_moe_layers(model)), []) + apply_pending_balance_updates(model) + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/tests/test_runtime/test_package_architecture.py b/tests/test_runtime/test_package_architecture.py index 9ede76de6..b0cfd681c 100644 --- a/tests/test_runtime/test_package_architecture.py +++ b/tests/test_runtime/test_package_architecture.py @@ -711,6 +711,7 @@ def test_dspark_configs_are_qwen3_family(self): self.assertEqual( set(dspark_configs), { + "deepseek-v4-flash-dspark-moe.json", "deepseek-v4-flash-dspark.json", "glm-5.2-dspark.json", "inkling-dspark.json",