From 379f2229133f6f26bb2a2071eb4995164aae56d7 Mon Sep 17 00:00:00 2001 From: morluto Date: Tue, 30 Jun 2026 20:46:33 +0000 Subject: [PATCH 1/4] Use float32 softmax for draft sampling --- deepspec/utils/sampling.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/deepspec/utils/sampling.py b/deepspec/utils/sampling.py index f395139b..4391c279 100644 --- a/deepspec/utils/sampling.py +++ b/deepspec/utils/sampling.py @@ -22,7 +22,7 @@ def sample_tokens(logits: torch.Tensor, temperature: float = 0.0) -> torch.Tenso return torch.argmax(logits, dim=-1) bsz, seq_len, vocab_size = logits.shape - flat_logits = logits.reshape(-1, vocab_size) / temperature + flat_logits = logits.reshape(-1, vocab_size).float() / temperature probs = torch.softmax(flat_logits, dim=-1) return torch.multinomial(probs, num_samples=1).reshape(bsz, seq_len) From f828b04da5cee5cee9a5ce6091c41624e9dbc646 Mon Sep 17 00:00:00 2001 From: morluto Date: Tue, 30 Jun 2026 20:47:34 +0000 Subject: [PATCH 2/4] Start cosine schedule from warmup peak --- deepspec/utils/optim.py | 1 + 1 file changed, 1 insertion(+) diff --git a/deepspec/utils/optim.py b/deepspec/utils/optim.py index 5e4341eb..032e6403 100644 --- a/deepspec/utils/optim.py +++ b/deepspec/utils/optim.py @@ -45,6 +45,7 @@ def get_lr(self): if not self.finished: self.after_scheduler.base_lrs = self.base_lrs self.finished = True + return self.base_lrs return self.after_scheduler.get_lr() return [(self.last_epoch + 1) / self.warmup_epochs * lr for lr in self.base_lrs] From bcd5d41b60eb59aedf6851f75a76bdf6dc223814 Mon Sep 17 00:00:00 2001 From: morluto Date: Tue, 30 Jun 2026 20:48:43 +0000 Subject: [PATCH 3/4] Pad Eagle3 fallback masks to cached length --- deepspec/modeling/eagle3/common.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/deepspec/modeling/eagle3/common.py b/deepspec/modeling/eagle3/common.py index a02e3d5b..0857dd39 100644 --- a/deepspec/modeling/eagle3/common.py +++ b/deepspec/modeling/eagle3/common.py @@ -162,6 +162,10 @@ def prepare_4d_causal_attention_mask( ) causal = causal.view(1, 1, q_len, kv_len) + mask_len = int(attention_mask.shape[-1]) + if mask_len < kv_len: + repeat_count = (kv_len + mask_len - 1) // mask_len + attention_mask = attention_mask.repeat(1, repeat_count) expanded_mask = attention_mask[:, None, None, :kv_len].to(device=device).bool() padding = torch.where( expanded_mask, From 67d7bdd46b5d532f43695968ea99492a7dddb442 Mon Sep 17 00:00:00 2001 From: morluto Date: Tue, 30 Jun 2026 21:16:49 +0000 Subject: [PATCH 4/4] Log DSpark loss with global weighting --- deepspec/modeling/dspark/loss.py | 29 ++++++----------------------- 1 file changed, 6 insertions(+), 23 deletions(-) diff --git a/deepspec/modeling/dspark/loss.py b/deepspec/modeling/dspark/loss.py index 88dd56ba..0b8ca837 100644 --- a/deepspec/modeling/dspark/loss.py +++ b/deepspec/modeling/dspark/loss.py @@ -274,23 +274,6 @@ def compute_dspark_loss( l1_loss_alpha = float(l1_loss_alpha) confidence_head_alpha = float(confidence_head_alpha) - local_ce_loss = loss_terms["ce_loss_num"] / (loss_terms["ce_loss_den"] + 1e-6) - local_l1_loss = local_ce_loss.new_zeros(()) - if global_denominators["l1_loss_den"].item() > 0: - local_l1_loss = loss_terms["l1_loss_num"] / ( - loss_terms["l1_loss_den"] + 1e-6 - ) - local_confidence_loss = local_ce_loss.new_zeros(()) - if has_confidence: - local_confidence_loss = loss_terms["confidence_loss_num"] / ( - loss_terms["confidence_loss_den"] + 1e-6 - ) - local_loss = ( - ce_loss_alpha * local_ce_loss - + l1_loss_alpha * local_l1_loss - + confidence_head_alpha * local_confidence_loss - ) - add_metric( "ce_loss", loss_terms["ce_loss_num"], @@ -311,12 +294,6 @@ def compute_dspark_loss( den=loss_terms["confidence_loss_den"], tag="train", ) - add_metric( - "loss", - local_loss, - reduction="mean", - tag="train", - ) backward_loss = _build_loss( loss_terms=loss_terms, global_denominators=global_denominators, @@ -326,6 +303,12 @@ def compute_dspark_loss( has_confidence=has_confidence, world_size=world_size, ) + add_metric( + "loss", + backward_loss.detach(), + reduction="dp_mean", + tag="train", + ) return backward_loss