Skip to content
Open
Show file tree
Hide file tree
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
29 changes: 6 additions & 23 deletions deepspec/modeling/dspark/loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"],
Expand All @@ -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,
Expand All @@ -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


Expand Down
4 changes: 4 additions & 0 deletions deepspec/modeling/eagle3/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
1 change: 1 addition & 0 deletions deepspec/utils/optim.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
2 changes: 1 addition & 1 deletion deepspec/utils/sampling.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down