Download code/model.py from ukung/semantic-lite-2-decoder-smoke-test: direct link, hf CLI and curl.
- Browser
- Download file 9.38 kB
-
https://huggingface.co/ukung/semantic-lite-2-decoder-smoke-test/resolve/main/code/model.py
- Command line
-
hf download hf://ukung/semantic-lite-2-decoder-smoke-test/code/model.py
-
curl -L -o model.py https://huggingface.co/ukung/semantic-lite-2-decoder-smoke-test/resolve/main/code/model.py
9.38 kB
| """ | |
| 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) | |
| 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) | |
| 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) | |