# How DependencyAwareLayer Enables Conditional s2 Token Generation Based on s1 Tokens in Kronos

> Learn how Kronos DependencyAwareLayer uses s1 token embeddings to conditionally generate s2 tokens via cross-attention, enhancing contextual generation.

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

---

**The DependencyAwareLayer is a lightweight cross-attention module that uses s1 token embeddings as queries to conditionally guide the generation of s2 tokens through contextual cross-attention over transformer hidden states.**

In the Kronos neural audio codec architecture, the model generates two discrete token streams: `s1` (coarse) and `s2` (fine). The **DependencyAwareLayer** bridges these streams by enabling **conditional s2 token generation based on s1 tokens**, ensuring that fine-grained predictions explicitly attend to their coarse counterparts. This mechanism is implemented in [`model/module.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/module.py) and integrated into the main forward pass in [`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py).

## Architecture of the DependencyAwareLayer

Located at lines 46–62 in [`model/module.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/module.py), the `DependencyAwareLayer` class implements a multi-head cross-attention mechanism with Rotary Position Embeddings (RoPE).

### Cross-Attention Mechanism with RoPE

The layer receives two inputs: `hidden_states` (the transformer context of shape `[B, T, D]`) and `sibling_embed` (the embedding of s1 tokens). It instantiates **MultiHeadCrossAttentionWithRoPE** to treat the s1 embedding as the **query** while using the hidden states as **key** and **value**:

```python
class DependencyAwareLayer(nn.Module):
    def __init__(self, d_model, n_heads=4, attn_dropout_p=0.0, resid_dropout=0.0):
        super().__init__()
        self.cross_attn = MultiHeadCrossAttentionWithRoPE(d_model, n_heads,
                                                         attn_dropout_p, resid_dropout)
        self.norm = RMSNorm(d_model)

    def forward(self, hidden_states, sibling_embed, key_padding_mask=None):
        # hidden_states: [B, T, D]  (output of the main transformer)

        # sibling_embed : [B, T, D]  (embedding of s1 tokens)

        attn_out = self.cross_attn(
            query=sibling_embed,
            key=hidden_states,
            value=hidden_states,
            key_padding_mask=key_padding_mask,
        )
        return self.norm(hidden_states + attn_out)

```

### Residual Connection and Normalization

The attention output is added residually to the original `hidden_states` and passed through **RMSNorm**. This preserves the transformer context while infusing it with s1-specific cues, producing a conditioned representation ready for s2 prediction.

## Integration in the Kronos Forward Pass

The layer is invoked in [`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py) (lines 74–76) after the s1 logits are computed, allowing the model to branch into **conditional s2 token generation**.

### Teacher Forcing vs. Sampling for s1

During training or inference, the model first obtains s1 token embeddings through either teacher forcing (using ground truth `s1_targets`) or sampling (using the softmax distribution over s1 logits):

```python

# Obtain sibling (s1) embedding

if use_teacher_forcing:
    sibling_embed = self.embedding.emb_s1(s1_targets)
else:
    s1_probs = F.softmax(s1_logits.detach(), dim=-1)
    sample_s1_ids = torch.multinomial(
        s1_probs.view(-1, self.s1_vocab_size), 1
    ).view(s1_ids.shape)
    sibling_embed = self.embedding.emb_s1(sample_s1_ids)

```

### The Conditioning Pipeline

The s1 embeddings are passed to the `DependencyAwareLayer` along with the transformer context `x`. The output `x2` feeds into `head.cond_forward()` to produce s2 logits:

```python

# Condition on s1 → produce s2 representation

x2 = self.dep_layer(x, sibling_embed, key_padding_mask=padding_mask)
s2_logits = self.head.cond_forward(x2)

```

By treating s1 as the query, the cross-attention extracts global context most relevant to each specific s1 token, ensuring contextual alignment between the two streams.

## Decoupled Inference with decode_s2

For inference scenarios requiring explicit separation of s1 and s2 generation, the `decode_s2` method (lines 26–28 in [`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py)) exposes the same conditioning logic:

```python
sibling_embed = self.embedding.emb_s1(s1_ids)
x2 = self.dep_layer(context, sibling_embed, key_padding_mask=padding_mask)
s2_logits = self.head.cond_forward(x2)

```

This allows pre-computed transformer contexts from `decode_s1` to be reused, running the **DependencyAwareLayer** only when s2 generation is required.

## Practical Implementation

### Manual Conditioning Without Teacher Forcing

The following example demonstrates how to manually condition s2 generation on sampled s1 tokens using the `DependencyAwareLayer`:

```python
import torch
from model.kronos import Kronos
from model.module import HierarchicalEmbedding, DependencyAwareLayer

# Model hyper-parameters (example)

model = Kronos(
    s1_bits=8, s2_bits=8,
    n_layers=4, d_model=256,
    n_heads=4, ff_dim=1024,
    ffn_dropout_p=0.1, attn_dropout_p=0.1,
    resid_dropout_p=0.1, token_dropout_p=0.1,
    learn_te=False,
).eval()

# Dummy inputs

batch, seq_len = 2, 10
s1_ids = torch.randint(0, 2**8, (batch, seq_len))
s2_ids = torch.randint(0, 2**8, (batch, seq_len))

# Forward pass up to the transformer context

x = model.embedding([s1_ids, s2_ids])          # hierarchical embedding

x = model.token_drop(x)
for layer in model.transformer:
    x = layer(x)                               # transformer blocks

x = model.norm(x)                              # RMSNorm

# Sample a new s1 sequence (no teacher forcing)

s1_logits = model.head(x)
s1_probs = torch.softmax(s1_logits, dim=-1)
sample_s1 = torch.multinomial(s1_probs.view(-1, model.s1_vocab_size), 1)
sample_s1 = sample_s1.view(batch, seq_len)

# Embed sampled s1 and run the DependencyAwareLayer

sibling_embed = model.embedding.emb_s1(sample_s1)
x2 = model.dep_layer(x, sibling_embed)        # <-- conditioning step

# Produce s2 logits conditioned on s1

s2_logits = model.head.cond_forward(x2)

```

### Using the decode_s2 Helper Method

For standard inference, use the built-in `decode_s2` method which encapsulates the conditioning logic:

```python

# Assume we already obtained the transformer context from decode_s1

s1_logits, context = model.decode_s1(s1_ids, s2_ids, stamp=None, padding_mask=None)

# Conditioned s2 logits

s2_logits = model.decode_s2(context, s1_ids)   # internally runs DependencyAwareLayer

```

## Summary

- The **DependencyAwareLayer** implements **conditional s2 token generation** via cross-attention where s1 embeddings serve as queries over transformer hidden states.
- Located in [`model/module.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/module.py) (lines 46–62), it uses **MultiHeadCrossAttentionWithRoPE** and **RMSNorm** to fuse s1 cues with global context.
- In [`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py) (lines 74–76), the layer processes either teacher-forced or sampled s1 embeddings to produce conditioned representations for the s2 prediction head.
- The `decode_s2` method (lines 26–28) provides decoupled inference, enabling separate generation of s1 and s2 token streams while maintaining dependency awareness.

## Frequently Asked Questions

### What is the role of sibling_embed in DependencyAwareLayer?

The `sibling_embed` parameter receives the embedding of s1 tokens and functions as the **query** tensor in the cross-attention mechanism. By querying the transformer hidden states (keys and values) with s1 embeddings, the layer extracts contextual information specifically relevant to each coarse token, enabling precise conditioning of the fine token stream.

### How does Kronos handle s1 token sampling during inference?

When teacher forcing is disabled, Kronos applies `F.softmax` to the s1 logits, samples token IDs using `torch.multinomial`, and looks up their embeddings via `self.embedding.emb_s1`. These sampled embeddings are then passed to the **DependencyAwareLayer** to condition s2 generation, as implemented in [`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py) lines 74–76.

### Why is cross-attention used instead of simple concatenation for conditioning s2 on s1?

Cross-attention allows the model to dynamically weigh the importance of different positions in the transformer context relative to each specific s1 token. Unlike concatenation, which treats all positions equally, the **MultiHeadCrossAttentionWithRoPE** mechanism in the **DependencyAwareLayer** enables fine-grained, position-aware fusion of s1 cues with the global hidden states, resulting in more accurate **conditional s2 token generation**.

### Can DependencyAwareLayer operate on pre-computed transformer contexts?

Yes. The `decode_s2` method in [`model/kronos.py`](https://github.com/shiyu-coder/Kronos/blob/main/model/kronos.py) (lines 26–28) demonstrates that the layer can process pre-computed contexts from `decode_s1`. By accepting a `context` argument and `s1_ids`, it enables decoupled inference where the expensive transformer computation happens once, and the lightweight **DependencyAwareLayer** conditions the output for s2 generation as needed.