# How Token Teacher Forcing Works in Kronos: Training and Inference Guide

> Understand token teacher forcing in Kronos training and inference. Learn how the use_teacher_forcing flag controls ground-truth vs sampled embeddings for efficient model operation.

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

---

**In the Kronos hierarchical decoder, token teacher forcing determines whether the s2 stage receives ground-truth s1 embeddings (training) or sampled s1 predictions (inference), switching via the `use_teacher_forcing` flag in [`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py).**

The `shiyu-coder/Kronos` repository implements a two-stage hierarchical transformer where the generation of s2 tokens depends on the s1 stage. Token teacher forcing controls how the **Dependency-Aware Layer** receives s1 context, allowing the model to train with stable ground-truth conditioning while inferring autoregressively from its own predictions.

## The `forward` Method and Token Teacher Forcing Logic

The core mechanism resides in the `forward` method of [`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py) (lines 239–276). This method accepts a boolean flag `use_teacher_forcing` and an optional tensor `s1_targets` containing the ground-truth s1 token IDs [[source](https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L239-L246)].

After computing the s1 logits, the model selects the embedding to feed into the dependency layer based on the forcing mode:

- **When `use_teacher_forcing=True`**: The model retrieves embeddings for the true s1 tokens via `self.embedding.emb_s1(s1_targets)` [[source](https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L267-L269)].
- **When `use_teacher_forcing=False`** (default): The model converts s1 logits to probabilities, samples token IDs using `torch.multinomial`, then looks up embeddings via `self.embedding.emb_s1(sample_s1_ids)` [[source](https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L270-L272)].

The selected embedding (`sibling_embed`) is concatenated with the transformer output `x` and passed to `self.dep_layer`, which injects s1 information into the representation for s2 decoding [[source](https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L274-L276)].

## Training vs. Inference Behavior

The token teacher forcing mechanism creates a clear separation between optimization and generation phases.

### Training Mode

During training, you explicitly set `use_teacher_forcing=True` and provide `s1_targets`. This forces the s2 decoder to observe the correct s1 context regardless of the model's current s1 predictions, stabilizing gradient flow and allowing independent loss computation for both stages.

### Inference Mode

During inference, the default `use_teacher_forcing=False` triggers autoregressive sampling. The model generates s1 tokens from its own probability distribution, then conditions s2 generation on these sampled tokens. For specialized inference pipelines, the repository also exposes `decode_s1` and `decode_s2` methods (lines 778–828) to separate the two stages manually.

## Code Examples

### Training with Teacher Forcing

The following snippet demonstrates a training loop that leverages ground-truth s1 tokens to condition s2 generation:

```python
import torch
from model.kronos import Kronos

# Initialize model

model = Kronos(s1_bits=8, s2_bits=8, n_layers=6,
              d_model=256, n_heads=8, ff_dim=1024,
              ffn_dropout_p=0.1, attn_dropout_p=0.1,
              resid_dropout_p=0.1, token_dropout_p=0.05,
              learn_te=True)

# Example batch

batch, seq_len = 32, 64
s1_ids = torch.randint(0, 2**8, (batch, seq_len))
s2_ids = torch.randint(0, 2**8, (batch, seq_len))
s1_targets = s1_ids.clone()
stamp = torch.arange(seq_len).unsqueeze(0).repeat(batch, 1)

# Forward with teacher forcing

s1_logits, s2_logits = model.forward(
    s1_ids, s2_ids,
    stamp=stamp,
    use_teacher_forcing=True,
    s1_targets=s1_targets)

# Compute separate losses

criterion = torch.nn.CrossEntropyLoss()
loss_s1 = criterion(s1_logits.view(-1, model.s1_vocab_size), s1_targets.view(-1))
loss_s2 = criterion(s2_logits.view(-1, model.s2_vocab_size), s2_ids.view(-1))
loss = loss_s1 + loss_s2
loss.backward()

```

### Inference Without Teacher Forcing

For generation, omit the `s1_targets` and allow the model to sample:

```python
model.eval()
with torch.no_grad():
    s1_ids = torch.zeros(batch, seq_len, dtype=torch.long)
    s2_ids = torch.zeros(batch, seq_len, dtype=torch.long)
    
    s1_logits, s2_logits = model.forward(
        s1_ids, s2_ids,
        stamp=stamp,
        use_teacher_forcing=False)  # Default behavior

    
    # Greedy decoding (or apply temperature sampling)

    s1_pred = torch.argmax(s1_logits, dim=-1)
    s2_pred = torch.argmax(s2_logits, dim=-1)

```

### Manual Two-Step Decoding

For beam search or custom sampling strategies, use the separate decoding methods:

```python

# Step 1: Generate s1

s1_logits, context = model.decode_s1(s1_ids, s2_ids, stamp=stamp)
s1_ids_pred = torch.argmax(s1_logits, dim=-1)

# Step 2: Generate s2 conditioned on s1

s2_logits = model.decode_s2(context, s1_ids_pred)
s2_ids_pred = torch.argmax(s2_logits, dim=-1)

```

## Key Implementation Files

- **[`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py)** (lines 239–276): Contains the `forward` method logic for `use_teacher_forcing` and embedding selection.
- **[`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py)** (lines 778–828): Implements `decode_s1` and `decode_s2` for stage-separated inference.
- **[`model/module.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/module.py)**: Provides the `HierarchicalEmbedding` class with `emb_s1` lookups and the `DependencyAwareLayer` that consumes the selected embeddings.

## Summary

- **Token teacher forcing** only affects the s1-to-s2 conditioning path, not the s1 prediction itself.
- Set `use_teacher_forcing=True` with `s1_targets` supplied during training to stabilize learning.
- Inference defaults to `False`, enabling the model to sample s1 tokens and condition s2 on its own predictions.
- The mechanism is implemented between lines 267–276 of [`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py), using `torch.multinomial` for sampling and `self.embedding.emb_s1` for embedding retrieval.

## Frequently Asked Questions

### Does token teacher forcing affect the s1 loss calculation?

No. The `use_teacher_forcing` flag only determines which s1 embeddings are passed to the **Dependency-Aware Layer** for s2 generation. The s1 logits are computed identically regardless of the flag, so the cross-entropy loss for s1 uses the same prediction distribution in both modes.

### What happens if I set `use_teacher_forcing=True` but forget to provide `s1_targets`?

The code at lines 267–269 expects `s1_targets` to be non-None when teacher forcing is enabled. If you omit this tensor, the model will raise a runtime error when attempting to call `self.embedding.emb_s1(s1_targets)`, as it cannot access embeddings for undefined ground-truth IDs.

### Can I mix teacher forcing and sampling within the same batch?

The current implementation uses a single boolean flag for the entire forward pass, meaning all sequences in the batch must use the same mode. To achieve mixed-mode training, you would need to run separate forward passes or modify the `forward` method to accept a per-sequence mask.

### How does the Dependency-AwareLayer use the s1 embedding?

According to lines 274–276 of [`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py), the layer receives both the transformer output `x` and the selected s1 embedding (`sibling_embed`). It fuses these representations to create a context vector that conditions the s2 prediction head, effectively allowing s2 tokens to attend to whichever s1 tokens were selected (ground-truth or sampled).