"""Nova v3 for transformers. Works on transformers 4.x and 5.x. The two versions disagree about how model code talks to the key/value cache: * 4.x passes the legacy format -- a tuple with one ``(keys, values)`` pair per layer -- or a ``Cache`` object, depending on version and settings; * 5.x always passes a ``Cache`` object (``DynamicCache`` and friends) and removed both ``Cache.__getitem__`` and the legacy tuples, so indexing ``past_key_values[0][0]`` raises ``TypeError``. This file writes to the cache through its ``update(keys, values, layer_idx)`` method, which both formats offer, and keeps the legacy tuple as a thin adapter. It also declares ``_tied_weights_keys`` as the ``{target: source}`` dict 5.x expects: Nova has a single already shared embedding matrix and ``get_output_embeddings()`` returns ``None``, so there is nothing to tie. """ import torch import torch.nn as nn import torch.nn.functional as F from transformers import GenerationMixin, PreTrainedModel from transformers.modeling_outputs import CausalLMOutputWithPast try: from .configuration_nova import NovaConfig except ImportError: from configuration_nova import NovaConfig def _rope(x, cos, sin): a, b = x.chunk(2, dim=-1) return torch.cat((a * cos - b * sin, b * cos + a * sin), dim=-1) class NovaRMSNorm(nn.Module): def __init__(self, dim): super().__init__() self.w = nn.Parameter(torch.ones(dim)) def forward(self, x): return F.rms_norm(x, (x.size(-1),), self.w, 1e-6) class _TupleCache: """The legacy cache of transformers 4.x: one ``(keys, values)`` pair per layer. It presents the same ``update`` method as ``transformers.cache_utils.Cache``, so blocks do not have to care which of the two they were handed, and can be turned back into the tuple format for callers that expect it. """ def __init__(self, layers=()): self.layers = list(layers) def update(self, keys, values, layer_idx): while len(self.layers) <= layer_idx: self.layers.append(None) past = self.layers[layer_idx] if past is not None and past[0] is not None and past[0].numel() > 0: keys = torch.cat((past[0], keys), dim=-2) values = torch.cat((past[1], values), dim=-2) self.layers[layer_idx] = (keys, values) return keys, values def to_tuple(self): return tuple(self.layers) def _as_cache(past_key_values): """Make both cache formats look alike: anything with an ``update`` method is used as is.""" if past_key_values is None or hasattr(past_key_values, "update"): return past_key_values return _TupleCache(past_key_values) def _cache_length(past_key_values): """How many positions are already stored in the cache.""" if past_key_values is None: return 0 if isinstance(past_key_values, _TupleCache): for pair in past_key_values.layers: if pair is not None and pair[0] is not None and pair[0].numel() > 0: return pair[0].size(-2) return 0 return past_key_values.get_seq_length() def _keep_mask(attention_mask, key_length): """Padding mask as booleans of shape ``(batch, 1, 1, key_length)``.""" if attention_mask.dim() == 4: # transformers 5 builds a 4D mask when a static cache is used. It is boolean for # sdpa, and additive (0 where the position is kept) in the older float form. mask = attention_mask if attention_mask.dtype == torch.bool else attention_mask == 0 else: mask = (attention_mask != 0)[:, None, None, :] return mask[..., -key_length:] class NovaBlock(nn.Module): def __init__(self, config): super().__init__() d, h, kv, hd, ff = ( config.hidden_size, config.num_attention_heads, config.num_key_value_heads, config.head_dim, config.intermediate_size, ) self.h, self.kv, self.hd = h, kv, hd self.n1, self.n2 = NovaRMSNorm(d), NovaRMSNorm(d) self.qn, self.kn = NovaRMSNorm(hd), NovaRMSNorm(hd) self.wq = nn.Linear(d, h * hd, bias=False) self.wk = nn.Linear(d, kv * hd, bias=False) self.wv = nn.Linear(d, kv * hd, bias=False) self.wo = nn.Linear(h * hd, d, bias=False) self.gate = nn.Linear(d, ff, bias=False) self.up = nn.Linear(d, ff, bias=False) self.down = nn.Linear(ff, d, bias=False) def forward(self, x, cos, sin, attention_mask, cache=None, layer_idx=0): batch, length, _ = x.shape h = self.n1(x) q = self.wq(h).view(batch, length, self.h, self.hd) k = self.wk(h).view(batch, length, self.kv, self.hd) v = self.wv(h).view(batch, length, self.kv, self.hd) q = _rope(self.qn(q), cos, sin).to(v.dtype).transpose(1, 2) k = _rope(self.kn(k), cos, sin).to(v.dtype).transpose(1, 2) v = v.transpose(1, 2) if cache is not None: # Store this pass and get back everything cached for this layer. k, v = cache.update(k, v, layer_idx) repeat = self.h // self.kv k = k.repeat_interleave(repeat, dim=1) v = v.repeat_interleave(repeat, dim=1) key_length = k.size(-2) past_length = key_length - length query_positions = past_length + torch.arange(length, device=x.device) key_positions = torch.arange(key_length, device=x.device) causal = key_positions[None, :] <= query_positions[:, None] mask = causal[None, None, :, :] if attention_mask is not None: mask = mask & _keep_mask(attention_mask, key_length) out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask) x = x + self.wo(out.transpose(1, 2).reshape(batch, length, -1)) h = self.n2(x) x = x + self.down(F.silu(self.gate(h)) * self.up(h)) return x class NovaModel(PreTrainedModel): config_class = NovaConfig base_model_prefix = "model" _no_split_modules = ["NovaBlock"] def __init__(self, config): super().__init__(config) self.emb = nn.Embedding(config.vocab_size, config.hidden_size) self.blocks = nn.ModuleList(NovaBlock(config) for _ in range(config.num_hidden_layers)) self.norm = NovaRMSNorm(config.hidden_size) self.rebuild_rope_buffers() self.post_init() def rebuild_rope_buffers(self): """(Re)build the rotary tables from the config. They are derived from the config, so they are not part of the checkpoint. That is fine on transformers 4, which runs ``__init__`` on the real device, but transformers 5 builds the model on the meta device and moves such buffers back uninitialized -- the tables would be left holding empty memory. It recomputes them by calling ``_init_weights``, which is why every use of the tables goes through this method. """ config = self.config frequencies = torch.outer( torch.arange(config.max_position_embeddings, dtype=torch.float32), 1 / config.rope_theta ** ( torch.arange(0, config.head_dim, 2, dtype=torch.float32) / config.head_dim ), ) self.register_buffer("cos", frequencies.cos(), persistent=False) self.register_buffer("sin", frequencies.sin(), persistent=False) @torch.no_grad() def _init_weights(self, module): if isinstance(module, NovaModel): module.rebuild_rope_buffers() else: super()._init_weights(module) def get_input_embeddings(self): return self.emb def set_input_embeddings(self, value): self.emb = value def forward( self, input_ids=None, attention_mask=None, position_ids=None, inputs_embeds=None, past_key_values=None, use_cache=None, **kwargs, ): if inputs_embeds is None: inputs_embeds = self.emb(input_ids) batch, length, _ = inputs_embeds.shape if use_cache is None: use_cache = getattr(self.config, "use_cache", False) cache = _as_cache(past_key_values) if use_cache else None if position_ids is None: if attention_mask is not None and attention_mask.dim() == 2: position_ids = attention_mask.long().cumsum(-1) - 1 position_ids.masked_fill_(attention_mask == 0, 0) position_ids = position_ids[:, -length:] else: past_length = _cache_length(cache) position_ids = torch.arange( past_length, past_length + length, device=inputs_embeds.device ).unsqueeze(0).expand(batch, -1) cos = self.cos[position_ids].unsqueeze(-2) sin = self.sin[position_ids].unsqueeze(-2) hidden = inputs_embeds for layer_idx, block in enumerate(self.blocks): hidden = block(hidden, cos, sin, attention_mask, cache, layer_idx) presents = None if use_cache: presents = cache if not isinstance(cache, _TupleCache) else cache.to_tuple() return self.norm(hidden), presents class NovaForCausalLM(PreTrainedModel, GenerationMixin): config_class = NovaConfig base_model_prefix = "model" # Nothing is tied: the embedding matrix is used for both the input and the output, and # get_output_embeddings() returns None. transformers 5 wants the mapping form, not a list. _tied_weights_keys = {} #: This model understands ``Cache`` objects, and also still accepts the legacy tuples. _supports_cache_class = True def __init__(self, config): super().__init__(config) self.model = NovaModel(config) self.post_init() def get_input_embeddings(self): return self.model.emb def set_input_embeddings(self, value): self.model.set_input_embeddings(value) def get_output_embeddings(self): return None @torch.no_grad() def _init_weights(self, module): # Recomputes the rotary tables, which are not in the checkpoint. See NovaModel. if isinstance(module, NovaModel): module.rebuild_rope_buffers() else: super()._init_weights(module) def prepare_inputs_for_generation( self, input_ids, attention_mask=None, token_type_ids=None, next_sequence_length=None, **kwargs, ): past_key_values = kwargs.get("past_key_values") past_length = _cache_length(past_key_values) if next_sequence_length is not None: # transformers 5 passes the full sequence and says how much of it is new. input_ids = input_ids[:, -next_sequence_length:] elif past_length > 0: # transformers 4 passes the full sequence as well, without saying how much is new. input_ids = input_ids[:, past_length:] if attention_mask is not None and attention_mask.dim() == 2: position_ids = attention_mask.long().cumsum(-1) - 1 position_ids.masked_fill_(attention_mask == 0, 0) position_ids = position_ids[:, -input_ids.size(1):] else: position_ids = torch.arange( past_length, past_length + input_ids.size(1), device=input_ids.device ).unsqueeze(0).expand(input_ids.size(0), -1) return { "input_ids": input_ids, "attention_mask": attention_mask, "position_ids": position_ids, "past_key_values": past_key_values, "use_cache": kwargs.get("use_cache", True), } def forward( self, input_ids=None, attention_mask=None, position_ids=None, inputs_embeds=None, labels=None, token_type_ids=None, past_key_values=None, use_cache=None, **kwargs, ): if use_cache is None: # transformers 5 moved `use_cache` out of the config, so it may be absent. use_cache = getattr(self.config, "use_cache", True) hidden, presents = self.model( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, inputs_embeds=inputs_embeds, past_key_values=past_key_values, use_cache=use_cache, ) logits = F.linear(hidden, self.model.emb.weight) loss = None if labels is not None: loss = F.cross_entropy( logits[..., :-1, :].contiguous().view(-1, logits.size(-1)), labels[..., 1:].contiguous().view(-1), ignore_index=-100, ) return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=presents )