nova-v3 / modeling_nova.py
plasmova's picture
Make the custom modeling code work on transformers 5: write to the cache through the Cache API, declare _tied_weights_keys as a mapping, and rebuild the rotary buffers after the meta-device load
501a7dd verified
Raw History Blame Contribute Delete
12.9 kB
"""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
)