How DependencyAwareLayer Enables Conditional s2 Token Generation Based on s1 Tokens in Kronos
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 and integrated into the main forward pass in model/kronos.py.
Architecture of the DependencyAwareLayer
Located at lines 46–62 in 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:
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 (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):
# 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:
# 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) exposes the same conditioning logic:
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:
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:
# 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(lines 46–62), it uses MultiHeadCrossAttentionWithRoPE and RMSNorm to fuse s1 cues with global context. - In
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_s2method (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 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 (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.
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →