How Kronos Computes Loss with DualHead for s1 and s2 Token Prediction
Kronos calculates the dual-head training loss by averaging two independent cross-entropy terms—one for the high-order s₁ token stream and one for the low-order s₂ token stream—while optionally masking padded positions to ensure only valid timesteps contribute to gradient updates.
The DualHead module in the shiyu-coder/Kronos repository manages parallel prediction heads for hierarchical time-series tokenization, where each continuous value is decomposed into coarse (s₁) and fine (s₂) bit representations. Understanding how compute_loss balances these two prediction tasks is essential for interpreting the model's training dynamics and diagnostic outputs.
DualHead Architecture and Forward Pass
In model/module.py, the DualHead class initializes two independent linear projections that transform the shared Transformer hidden state into separate vocabulary distributions:
# model/module.py
class DualHead(nn.Module):
def __init__(self, s1_bits, s2_bits, d_model):
super().__init__()
self.vocab_s1 = 2 ** s1_bits # size of s₁ vocabulary
self.vocab_s2 = 2 ** s2_bits # size of s₂ vocabulary
self.proj_s1 = nn.Linear(d_model, self.vocab_s1) # s₁ head
self.proj_s2 = nn.Linear(d_model, self.vocab_s2) # s₂ head
During the forward pass through model/kronos.py, the model generates logits for both streams:
- s₁ logits: Produced by
self.head(x)for the high-order bits (pre-token). - s₂ logits: Produced by
self.head.cond_forward(x2)after injecting the already-predicted s₁ context through a dependency-aware layer.
This architectural choice ensures that s₂ prediction conditions on the s₁ representation, creating a hierarchical refinement pattern.
The compute_loss Method Implementation
The compute_loss method (lines 94–107 in model/module.py) handles the mathematical combination of both prediction errors.
Cross-Entropy Calculation and Masking
The method accepts s1_logits, s2_logits, s1_targets, s2_targets, and an optional padding_mask. When a mask is provided, the function filters out padded positions before computing loss:
# model/module.py (lines 94-107)
def compute_loss(self, s1_logits, s2_logits, s1_targets, s2_targets,
padding_mask=None):
if padding_mask is not None:
# ignore padded positions
valid_mask = (padding_mask == 0)
s1_logits = s1_logits[valid_mask]
s2_logits = s2_logits[valid_mask]
s1_targets = s1_targets[valid_mask]
s2_targets = s2_targets[valid_mask]
ce_s1 = F.cross_entropy(s1_logits, s1_targets)
ce_s2 = F.cross_entropy(s2_logits, s2_targets)
else:
# flatten batch-seq dimensions
ce_s1 = F.cross_entropy(
s1_logits.reshape(-1, self.vocab_s1), s1_targets.reshape(-1))
ce_s2 = F.cross_entropy(
s2_logits.reshape(-1, self.vocab_s2), s2_targets.reshape(-1))
# final loss = average of the two CE terms
ce_loss = (ce_s1 + ce_s2) / 2
return ce_loss, ce_s1, ce_s2
If no padding mask is supplied, the implementation reshapes the batch and sequence dimensions into a single dimension using .reshape(-1, ...) to compute standard cross-entropy across all positions.
Loss Aggregation Strategy
The method returns three tensors: the averaged loss (ce_loss) used for backpropagation, and the individual s₁ (ce_s1) and s₂ (ce_s2) components for monitoring. This equal weighting ensures that neither the coarse nor fine representation dominates the gradient updates during training.
Training Loop Integration
In practice, the loss computation is invoked within the training script at finetune_csv/finetune_base_model.py (lines 91–94):
# finetune_csv/finetune_base_model.py (lines 91-94)
logits = (model.module if use_ddp else model)(
token_in[0], token_in[1], batch_x_stamp[:, :-1, :])
loss, s1_loss, s2_loss = (model.module if use_ddp else model).head.compute_loss(
logits[0], logits[1], token_out[0], token_out[1])
Here, token_in contains the teacher-forced inputs (shifted right by one position), while token_out holds the target s₁ and s₂ token streams. The model returns a tuple of logits that passes directly into compute_loss, which then drives the optimizer step.
Practical Code Examples
Example 1: Standalone Loss Computation
Use the DualHead module directly to compute gradients on synthetic data:
import torch
import torch.nn.functional as F
from model.module import DualHead
# Configuration
batch, seq_len, d_model = 4, 32, 256
s1_bits, s2_bits = 8, 8
# Synthetic forward pass outputs
s1_logits = torch.randn(batch, seq_len, 2 ** s1_bits, requires_grad=True)
s2_logits = torch.randn(batch, seq_len, 2 ** s2_bits, requires_grad=True)
s1_targets = torch.randint(0, 2 ** s1_bits, (batch, seq_len))
s2_targets = torch.randint(0, 2 ** s2_bits, (batch, seq_len))
# Initialize head and compute loss
head = DualHead(s1_bits, s2_bits, d_model)
loss, ce_s1, ce_s2 = head.compute_loss(s1_logits, s2_logits,
s1_targets, s2_targets)
loss.backward() # Gradients flow to both projection layers
print(f"Total: {loss.item():.4f} (s1={ce_s1.item():.4f}, s2={ce_s2.item():.4f})")
Example 2: Training Loop Implementation
Integrate the dual-head loss into a standard PyTorch training iteration:
for batch_x, batch_stamp in train_loader:
# Tokenize into hierarchical streams
token_pre, token_post = tokenizer.encode(batch_x, half=True)
# Prepare teacher-forcing inputs and targets
token_in = [token_pre[:, :-1], token_post[:, :-1]]
token_out = [token_pre[:, 1:], token_post[:, 1:]]
# Forward pass through Kronos
s1_logits, s2_logits = model(token_in[0], token_in[1], batch_stamp[:, :-1, :])
# Compute dual-head loss
loss, s1_ce, s2_ce = model.head.compute_loss(
s1_logits, s2_logits, token_out[0], token_out[1])
optimizer.zero_grad()
loss.backward()
optimizer.step()
Summary
- DualHead architecture: Uses separate
proj_s1andproj_s2linear layers inmodel/module.pyto predict high-order and low-order bit tokens from shared Transformer states. - Loss computation: The
compute_lossmethod averages two cross-entropy terms, ensuring balanced optimization of both s₁ and s₂ distributions. - Masking support: When provided,
padding_maskfilters invalid positions, computing loss only on valid sequence elements. - Diagnostic outputs: The method returns individual
ce_s1andce_s2values alongside the combined loss for training monitoring. - Training integration: Called directly in
finetune_csv/finetune_base_model.pywith logits from the dependency-aware forward pass.
Frequently Asked Questions
How does Kronos handle variable-length sequences in the dual-head loss?
When a padding_mask is provided to compute_loss, the method creates a valid_mask where positions equal to zero are considered valid. It then indexes all logits and targets with this boolean mask before computing cross-entropy, ensuring that padded timesteps contribute nothing to the gradient update.
Why does Kronos use an average instead of a weighted sum for the two cross-entropy terms?
The implementation computes ce_loss = (ce_s1 + ce_s2) / 2 to treat both the coarse s₁ prediction and the fine s₂ refinement as equally important learning objectives. This prevents either head from dominating the optimization landscape and maintains balanced representation learning across both bit levels.
What is the difference between s₁ and s₂ tokens in the Kronos architecture?
s₁ tokens represent the high-order bits (pre-tokens) of the composite code, providing coarse quantization of the time-series values, while s₂ tokens represent the low-order bits (post-tokens) that refine the representation. The s₂ prediction depends on the s₁ output through the cond_forward method, creating a hierarchical generation pattern.
Where is the compute_loss method defined in the Kronos codebase?
The compute_loss method is defined in model/module.py at lines 94–107 within the DualHead class. This module is instantiated as self.head in the main Kronos model and invoked during training in finetune_csv/finetune_base_model.py.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →