# How Kronos Computes Loss with DualHead for s1 and s2 Token Prediction

> Learn how Kronos computes loss with DualHead for s1 and s2 token prediction by averaging cross-entropy terms and optionally masking padded positions for accurate gradient updates.

- Repository: [ShiYu/Kronos](https://github.com/shiyu-coder/Kronos)
- Tags: deep-dive
- Published: 2026-04-10

---

**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`](https://github.com/shiyu-coder/Kronos/blob/main/model/module.py), the `DualHead` class initializes two independent linear projections that transform the shared Transformer hidden state into separate vocabulary distributions:

```python

# 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`](https://github.com/shiyu-coder/Kronos/blob/main/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`](https://github.com/shiyu-coder/Kronos/blob/main/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:

```python

# 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`](https://github.com/shiyu-coder/Kronos/blob/main/finetune_csv/finetune_base_model.py) (lines 91–94):

```python

# 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:

```python
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:

```python
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`](https://github.com/shiyu-coder/Kronos/blob/main/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`](https://github.com/shiyu-coder/Kronos/blob/main/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`](https://github.com/shiyu-coder/Kronos/blob/main/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`](https://github.com/shiyu-coder/Kronos/blob/main/finetune_csv/finetune_base_model.py).