""" Semantic-Conditioned Decoder. input text -> Semantic-Lite-2 (FROZEN) Data A: (B, 256) -> proj_a -> prefix token prepended to decoder Data B: (B, L, 2048) -> proj_b -> cross-attention key/value -> TransformerDecoderLayer (d_model=1024, nhead=16, FFN=1024) [TRAINABLE] -> output_proj (1024 -> 2048) + tied embedding (131072 vocab) Only proj_a, proj_b, dec_embed_proj, the decoder layer, output_proj and pos_embed are trained. The encoder and its embedding table are frozen. TIED EMBEDDING -------------- The output head reuses the encoder's frozen embedding matrix instead of learning a 131072 x 2048 matrix. That saves ~132M parameters (and the VRAM to hold them), at the cost of forcing the output geometry to match an embedding table that was never trained for generation. Whether that trade is worth it is an open question — see NOTES.md. """ import torch import torch.nn as nn import torch.nn.functional as F DATA_A_DIM = 256 DATA_B_DIM = 2048 class SemanticConditionedDecoder(nn.Module): def __init__(self, encoder, tokenizer, d_model=1024, nhead=16, num_decoder_layers=1, dim_feedforward=1024, dropout=0.1): super().__init__() self.encoder = encoder # frozen self.tokenizer = tokenizer self.d_model = d_model self.vocab_size = tokenizer.vocab_size self.hidden_size = encoder.config.hidden_size # 2048 self.proj_a = nn.Linear(DATA_A_DIM, d_model) self.proj_b = nn.Linear(DATA_B_DIM, d_model) self.dec_embed_proj = nn.Linear(self.hidden_size, d_model, bias=False) decoder_layer = nn.TransformerDecoderLayer( d_model=d_model, nhead=nhead, dim_feedforward=dim_feedforward, dropout=dropout, batch_first=True, activation="gelu", ) self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=num_decoder_layers) self.output_proj = nn.Linear(d_model, self.hidden_size, bias=False) self.pos_embed = nn.Embedding(2048, d_model) self.bos_id = tokenizer.bos_token_id self.eos_id = tokenizer.eos_token_id self.pad_id = tokenizer.pad_token_id def _logits_from_hidden(self, hidden): """hidden (..., d_model) -> logits (..., vocab_size) via tied embedding.""" h = self.output_proj(hidden) # (..., 2048) emb = self.encoder.backbone.embedding.weight # (131072, 2048), frozen return F.linear(h, emb) def encode(self, input_ids, attention_mask): """Encode input -> Data A (B, 256) + Data B (B, L, 2048).""" with torch.no_grad(): data_a = self.encoder(input_ids=input_ids, attention_mask=attention_mask) out_b = self.encoder.backbone(input_ids=input_ids, attention_mask=attention_mask) data_b = out_b.last_hidden_state return data_a, data_b def _embed_target(self, decoder_input_ids): """Token ids -> d_model embeddings + learned positional embedding.""" emb = self.encoder.backbone.embedding(decoder_input_ids) # (B, L, 2048) emb = self.dec_embed_proj(emb) # (B, L, d_model) pos = torch.arange(decoder_input_ids.shape[1], device=decoder_input_ids.device) return emb + self.pos_embed(pos).unsqueeze(0) @staticmethod def _causal_mask(length, device): """ Bool causal mask: True above the diagonal = "not allowed to attend". Bool (not float -inf) so that it matches the dtype of tgt_key_padding_mask. Mixing a float attn_mask with a bool key_padding_mask is deprecated in recent PyTorch and emits a warning. """ return torch.triu( torch.ones(length, length, dtype=torch.bool, device=device), diagonal=1 ) def forward(self, input_ids, attention_mask, decoder_input_ids): """ input_ids: (B, L_src) source text (the `problem` column) attention_mask: (B, L_src) decoder_input_ids: (B, L_tgt) shifted-right target (BOS + thinking + solution) """ B = input_ids.shape[0] data_a, data_b = self.encode(input_ids, attention_mask) a_proj = self.proj_a(data_a) # (B, d_model) b_proj = self.proj_b(data_b) # (B, L_src, d_model) dec_emb = self._embed_target(decoder_input_ids) # Data A as a prefix token, so the global meaning is visible at every step. dec_emb = torch.cat([a_proj.unsqueeze(1), dec_emb], dim=1) # (B, 1+L_tgt, d_model) L_dec = dec_emb.shape[1] tgt_mask = self._causal_mask(L_dec, dec_emb.device) # Prefix token is never padding, so prepend a False column. tgt_key_padding_mask = torch.cat([ torch.zeros(B, 1, dtype=torch.bool, device=decoder_input_ids.device), (decoder_input_ids == self.pad_id), ], dim=1) dec_out = self.decoder( tgt=dec_emb, memory=b_proj, tgt_mask=tgt_mask, tgt_key_padding_mask=tgt_key_padding_mask, memory_key_padding_mask=(attention_mask == 0), ) dec_out = dec_out[:, 1:, :] # drop the prefix token return self._logits_from_hidden(dec_out) @torch.no_grad() def generate(self, input_ids, attention_mask, max_new_tokens=256, temperature=1.0): """Autoregressive generation. Returns (B, 1 + n_generated) token ids.""" self.eval() B, device = input_ids.shape[0], input_ids.device data_a, data_b = self.encode(input_ids, attention_mask) a_proj = self.proj_a(data_a) memory = self.proj_b(data_b) mem_pad_mask = (attention_mask == 0) generated = torch.full((B, 1), self.bos_id, dtype=torch.long, device=device) for _ in range(max_new_tokens): dec_emb = self._embed_target(generated) dec_emb = torch.cat([a_proj.unsqueeze(1), dec_emb], dim=1) L_dec = dec_emb.shape[1] tgt_mask = self._causal_mask(L_dec, device) dec_out = self.decoder( tgt=dec_emb, memory=memory, tgt_mask=tgt_mask, memory_key_padding_mask=mem_pad_mask, ) logits = self._logits_from_hidden(dec_out[:, -1:, :]) if temperature != 1.0: logits = logits / temperature probs = F.softmax(logits, dim=-1) next_token = torch.multinomial(probs.squeeze(1), num_samples=1) generated = torch.cat([generated, next_token], dim=1) if (next_token == self.eos_id).all(): break return generated class DecoderNoConditioning(nn.Module): """ Ablation baseline: identical to SemanticConditionedDecoder except the cross-attention memory is a learned constant instead of a function of the input, and there is no Data A prefix. proj_a / proj_b are kept (but unused) so the trainable parameter count matches the conditioned arm exactly. The only variable that changes is whether the memory carries information about the input. """ def __init__(self, encoder, tokenizer, d_model=1024, nhead=16, num_decoder_layers=1, dim_feedforward=1024, dropout=0.1): super().__init__() self.encoder = encoder self.tokenizer = tokenizer self.d_model = d_model self.vocab_size = tokenizer.vocab_size self.hidden_size = encoder.config.hidden_size # Parameter parity only — never used in forward(). self.proj_a = nn.Linear(DATA_A_DIM, d_model) self.proj_b = nn.Linear(DATA_B_DIM, d_model) self.dec_embed_proj = nn.Linear(self.hidden_size, d_model, bias=False) decoder_layer = nn.TransformerDecoderLayer( d_model=d_model, nhead=nhead, dim_feedforward=dim_feedforward, dropout=dropout, batch_first=True, activation="gelu", ) self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=num_decoder_layers) self.output_proj = nn.Linear(d_model, self.hidden_size, bias=False) self.pos_embed = nn.Embedding(2048, d_model) self.null_memory = nn.Parameter(torch.randn(1, 1, d_model) * 0.02) self.bos_id = tokenizer.bos_token_id self.eos_id = tokenizer.eos_token_id self.pad_id = tokenizer.pad_token_id def _logits_from_hidden(self, hidden): h = self.output_proj(hidden) return F.linear(h, self.encoder.backbone.embedding.weight) def forward(self, input_ids, attention_mask, decoder_input_ids): B = input_ids.shape[0] memory = self.null_memory.expand(B, -1, -1) # (B, 1, d_model) emb = self.encoder.backbone.embedding(decoder_input_ids) emb = self.dec_embed_proj(emb) pos = torch.arange(decoder_input_ids.shape[1], device=decoder_input_ids.device) emb = emb + self.pos_embed(pos).unsqueeze(0) L_dec = emb.shape[1] tgt_mask = torch.triu( torch.ones(L_dec, L_dec, dtype=torch.bool, device=emb.device), diagonal=1 ) dec_out = self.decoder( tgt=emb, memory=memory, tgt_mask=tgt_mask, tgt_key_padding_mask=(decoder_input_ids == self.pad_id), ) return self._logits_from_hidden(dec_out)