diff --git a/experimental/builder/models/dspark/modeling_dspark_draft.py b/experimental/builder/models/dspark/modeling_dspark_draft.py index 41b0189ce..451af8421 100644 --- a/experimental/builder/models/dspark/modeling_dspark_draft.py +++ b/experimental/builder/models/dspark/modeling_dspark_draft.py @@ -18,11 +18,13 @@ import tensorrt as trt +from ...core import config as core_config from ...core import quantization from ...ops import (GatedMLP, Linear, Module, NetworkModule, RMSNorm, TreeAttention) from ...ops import functional as F from ...ops import pack_qkv +from .. import registry as model_registry class DSparkProposalAttention(TreeAttention): @@ -112,11 +114,45 @@ def forward(self, hidden): class DSparkDraftModel(NetworkModule): """Parallel proposal backbone with hidden output for sequential heads.""" - def __init__(self, ctx) -> None: - super().__init__(ctx) - if not (ctx.weights.has("lm_head.weight") + @classmethod + def from_config(cls, ctx): + if (ctx.weights.has("lm_head.weight") or ctx.weights.has("lm_head.qweight")): - raise ValueError("DSpark draft checkpoint must provide lm_head") + return cls(ctx) + + # Drafts such as RadixArk/Qwen3.8-27B-DSpark ship no head of their own + # and share the target model's lm_head, quantized as in the target. + args = ctx.args + target_cfg = core_config.DeviceConfig.from_pretrained( + args.target_model_dir, tp_size=args.tp_size, tp_rank=args.tp_rank) + target_bundle = core_config.BundleConfig.from_pretrained( + args.target_model_dir) + conversion = model_registry.weight_conversion_for( + target_bundle.root_model_type) + target_weights = ctx.open_weights( + args.target_model_dir, + group_size=target_cfg.group_size, + quant=target_cfg.quant, + component="llm", + vocab_map=ctx.weights.vocab_map, + conversion=conversion, + int4_gemm_plugin_version=args.int4_gemm_plugin_version, + checkpoint_source="target", + tie_word_embeddings=target_cfg.tie_word_embeddings) + try: + target_context = ctx.with_checkpoint(target_cfg, target_weights) + model = cls(ctx, + lm_head=Linear(target_context, + target_weights.causal_lm_head_prefix())) + except Exception: + target_weights.close() + raise + model._target_weights = target_weights + return model + + def __init__(self, ctx, lm_head=None) -> None: + super().__init__(ctx) + self._target_weights = None self.fc = DSparkTargetProjection(ctx, "fc") self.hidden_norm = RMSNorm(ctx, "hidden_norm", ctx.cfg.rms_norm_eps) self.layers = [ @@ -124,7 +160,13 @@ def __init__(self, ctx) -> None: for index in range(ctx.cfg.num_hidden_layers) ] self.norm = RMSNorm(ctx, "norm", ctx.cfg.rms_norm_eps) - self.lm_head = Linear(ctx, "lm_head") + self.lm_head = lm_head or Linear(ctx, "lm_head") + + def close(self) -> None: + if self._target_weights is not None: + self._target_weights.close() + self._target_weights = None + super().close() def input_tensors(self) -> Dict[str, object]: cfg = self.cfg