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

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.

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 (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].

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].
  • 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].

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].

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:

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:

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:


# 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 (lines 239–276): Contains the forward method logic for use_teacher_forcing and embedding selection.
  • model/kronos.py (lines 778–828): Implements decode_s1 and decode_s2 for stage-separated inference.
  • 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, 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, 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).

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 →