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_s1 and proj_s2 linear layers in model/module.py to predict high-order and low-order bit tokens from shared Transformer states.
  • Loss computation: The compute_loss method averages two cross-entropy terms, ensuring balanced optimization of both s₁ and s₂ distributions.
  • Masking support: When provided, padding_mask filters invalid positions, computing loss only on valid sequence elements.
  • Diagnostic outputs: The method returns individual ce_s1 and ce_s2 values alongside the combined loss for training monitoring.
  • Training integration: Called directly in finetune_csv/finetune_base_model.py with 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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →