Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 47 additions & 5 deletions experimental/builder/models/dspark/modeling_dspark_draft.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -112,19 +114,59 @@ 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 = [
DSparkDecoderLayer(ctx, f"layers.{index}")
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
Expand Down