LOMONOSOV-ZENIT-27B-1M-INDEV / modeling_altay.py
Ddavidich's picture
Кривая удержания переснята на 48 закладках; RazorAttention и вход transformers доехали до репозитория
554a918 verified
Raw
History Blame Contribute Delete
19.7 kB
"""ALTAY для transformers: `from_pretrained` без vLLM.
До появления этого файла путь через transformers не работал вовсе. Загрузчик
откатывался на стоковый `qwen3_5`, который не знает ни упакованных в восемь бит
эмбеддингов, ни оверлея ALTAY, и падал на `'Embedding' object has no attribute
'weight'`.
Здесь воспроизведены обе недостающие части, и воспроизведены **целиком**.
Соблазн собрать урезанную версию был велик — множители оверлея малы, от −0.014 до
+0.023, и казалось, что без него «почти то же самое». Так делать нельзя: под
кнопкой «Use this model» оказалась бы модель, отвечающая иначе, чем отгружаемая.
Либо честно, либо никак.
Что такое оверлей. Двадцать восемь тензоров, два суперблока. Суперблок — это три
рекуррентных повтора и один повтор внимания подряд. Повтор берёт **чужой слой
целиком** — его нормировку, смеситель и MLP, с теми же весами, — прогоняет через
него состояние ещё раз и добавляет результат с обучаемым множителем
`altay_alpha`. Сверх этого к запросу и выходу внимания добавляются низкоранговые
дельты ранга 64.
суперблок A: повторяет слои 28, 29, 30 (рекуррентные) и 31 (внимание),
выполняется сразу после слоя 31;
суперблок B: повторяет слои 60, 61, 62 и 63, выполняется после слоя 63.
Приватное состояние. Рекуррентные повторы обязаны иметь **своё** состояние, а не
делить его с исходным слоем: в vLLM для этого заведены шесть дополнительных
ячеек кэша. Здесь та же вещь сделана вторым объектом кэша — веса общие, состояния
раздельные.
Чего этот файл НЕ даёт по сравнению с путём через vLLM: кодека TurboQuant для
кэша KV, профилей запуска и ступенчатой карты позиций. На окнах до 196 608
токенов карта позиций тождественна, поэтому там разницы нет; дальше — есть, и
для длинного контекста надо брать vLLM.
"""
from __future__ import annotations
import torch
from torch import nn
from torch.nn import functional as F
from transformers.masking_utils import create_causal_mask
from transformers.models.qwen3_5.modeling_qwen3_5 import (
Qwen3_5ForConditionalGeneration,
Qwen3_5Model,
Qwen3_5ModelOutputWithPast,
Qwen3_5TextModel,
apply_rotary_pos_emb,
)
try: # имя менялось между версиями transformers
from transformers.masking_utils import create_recurrent_attention_mask
except ImportError: # pragma: no cover
create_recurrent_attention_mask = create_causal_mask
OVERLAY_RANK = 64
SUPERBLOCK_A_RECURRENT = (28, 29, 30)
SUPERBLOCK_A_ATTENTION = 31
SUPERBLOCK_B_RECURRENT = (60, 61, 62)
SUPERBLOCK_B_ATTENTION = 63
def _linear(in_features: int, out_features: int, reference: torch.Tensor) -> nn.Linear:
layer = nn.Linear(in_features, out_features, bias=False)
return layer.to(device=reference.device, dtype=reference.dtype)
class AltayRecurrentReplay(nn.Module):
"""Повтор рекуррентного слоя с приватным состоянием и дельтой ранга 64."""
def __init__(self, source_layer: nn.Module, hidden_size: int) -> None:
super().__init__()
# Слой-источник не регистрируется как подмодуль: его веса уже есть в
# чекпойнте по своему адресу, и второй адрес сломал бы загрузку.
object.__setattr__(self, "_source", source_layer)
reference = source_layer.input_layernorm.weight
self.altay_alpha = nn.Parameter(
torch.zeros((), device=reference.device, dtype=torch.float32)
)
self.delta_down = _linear(hidden_size, OVERLAY_RANK, reference)
self.delta_up = _linear(OVERLAY_RANK, hidden_size, reference)
@property
def source(self) -> nn.Module:
return object.__getattribute__(self, "_source")
def forward(self, hidden_states, *, attention_mask, replay_cache, **kwargs):
layer = self.source
normalized = layer.input_layernorm(hidden_states)
mixed = layer.linear_attn(
hidden_states=normalized,
cache_params=replay_cache,
attention_mask=attention_mask,
**kwargs,
)
transformed = hidden_states + mixed
transformed = transformed + layer.mlp(layer.post_attention_layernorm(transformed))
delta = self.delta_up(F.silu(self.delta_down(hidden_states)))
branch = transformed - hidden_states + delta
return hidden_states + self.altay_alpha.to(hidden_states.dtype) * branch
class AltayAttentionReplay(nn.Module):
"""Повтор слоя внимания: свежие запрос и выход поверх тех же ключей."""
def __init__(self, source_layer: nn.Module, hidden_size: int) -> None:
super().__init__()
object.__setattr__(self, "_source", source_layer)
attention = source_layer.self_attn
reference = source_layer.input_layernorm.weight
qg_width = attention.q_proj.out_features # запрос и вентиль вместе
o_input = attention.o_proj.in_features
self.altay_alpha = nn.Parameter(
torch.zeros((), device=reference.device, dtype=torch.float32)
)
self.q_delta_down = _linear(hidden_size, OVERLAY_RANK, reference)
self.q_delta_up = _linear(OVERLAY_RANK, qg_width, reference)
self.o_delta_down = _linear(o_input, OVERLAY_RANK, reference)
self.o_delta_up = _linear(OVERLAY_RANK, hidden_size, reference)
@property
def source(self) -> nn.Module:
return object.__getattribute__(self, "_source")
def forward(self, hidden_states, *, position_embeddings, attention_mask,
position_ids=None, **kwargs):
from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
from transformers.models.qwen3_5.modeling_qwen3_5 import eager_attention_forward
layer = self.source
attention = layer.self_attn
head_dim = attention.head_dim
normalized = layer.input_layernorm(hidden_states)
input_shape = normalized.shape[:-1]
hidden_shape = (*input_shape, -1, head_dim)
# Дельта ложится на запрос ДО нормировки головы и до поворота позиций —
# именно там она обучалась.
qg = attention.q_proj(normalized)
qg = qg + self.q_delta_up(F.silu(self.q_delta_down(normalized)))
query, gate = torch.chunk(qg.view(*input_shape, -1, head_dim * 2), 2, dim=-1)
gate = gate.reshape(*input_shape, -1)
query = attention.q_norm(query.view(hidden_shape)).transpose(1, 2)
key = attention.k_norm(attention.k_proj(normalized).view(hidden_shape)).transpose(1, 2)
value = attention.v_proj(normalized).view(hidden_shape).transpose(1, 2)
cos, sin = position_embeddings
query, key = apply_rotary_pos_emb(query, key, cos, sin)
# Кэш не трогаем: повтор считает по своим ключам и значениям и ничего
# в общий кэш не пишет, иначе он испортил бы основной проход.
interface = ALL_ATTENTION_FUNCTIONS.get_interface(
attention.config._attn_implementation, eager_attention_forward)
raw, _ = interface(attention, query, key, value, attention_mask,
dropout=0.0, scaling=attention.scaling)
raw = raw.reshape(*input_shape, -1).contiguous()
mixed = attention.o_proj(raw * torch.sigmoid(gate))
mixed = mixed + self.o_delta_up(F.silu(self.o_delta_down(raw)))
transformed = hidden_states + mixed
transformed = transformed + layer.mlp(layer.post_attention_layernorm(transformed))
branch = transformed - hidden_states
return hidden_states + self.altay_alpha.to(hidden_states.dtype) * branch
class AltayOverlay(nn.Module):
"""Двадцать восемь тензоров оверлея: два суперблока по четыре повтора."""
def __init__(self, text_model: Qwen3_5TextModel) -> None:
super().__init__()
hidden = text_model.config.hidden_size
self.sb_a_recurrent = nn.ModuleList(
[AltayRecurrentReplay(text_model.layers[i], hidden)
for i in SUPERBLOCK_A_RECURRENT])
self.sb_a_attention = AltayAttentionReplay(
text_model.layers[SUPERBLOCK_A_ATTENTION], hidden)
self.sb_b_recurrent = nn.ModuleList(
[AltayRecurrentReplay(text_model.layers[i], hidden)
for i in SUPERBLOCK_B_RECURRENT])
self.sb_b_attention = AltayAttentionReplay(
text_model.layers[SUPERBLOCK_B_ATTENTION], hidden)
def run(self, name: str, hidden_states, *, recurrent_mask, attention_mask,
position_embeddings, position_ids, replay_cache, **kwargs):
for replay in getattr(self, f"{name}_recurrent"):
hidden_states = replay(hidden_states, attention_mask=recurrent_mask,
replay_cache=replay_cache, **kwargs)
return getattr(self, f"{name}_attention")(
hidden_states, position_embeddings=position_embeddings,
attention_mask=attention_mask, position_ids=position_ids, **kwargs)
class LomonosovZenitAltayTextModel(Qwen3_5TextModel):
"""Стоковая текстовая модель плюс оверлей после слоёв 31 и 63."""
def __init__(self, config):
super().__init__(config)
self.altay_overlay = AltayOverlay(self)
# ---- упакованные в восемь бит эмбеддинги -------------------------------
def materialize_packed_embedding(self) -> None:
"""Развернуть embed_tokens ПОСЛЕ загрузки весов.
Прежде здесь стоял пре-хук `_register_load_state_dict_pre_hook`,
собиравший `weight` из пары `weight_packed`/`weight_scale` прямо в
словаре состояния. Замерено 26.07.2026, что он не работает и работать не
может по двум причинам сразу.
Во-первых, compressed-tensors 0.17 обрабатывает `nn.Embedding` наравне с
`nn.Linear`: при подготовке модуля он сам снимает `weight` и заводит
`weight_packed`, `weight_scale`, `weight_shape` как параметры. Значит
собирать `weight` из словаря состояния не нужно — упакованные тензоры
загружаются штатно, MISSING нет ни одного. Не хватает только обратного
хода, потому что сжатого прохода вперёд у эмбеддинга нет.
Во-вторых, хук обращался к `module.weight.dtype` — к тому самому
атрибуту, отсутствие которого он и должен был лечить.
Поэтому распаковываем после загрузки и распаковщиком самой библиотеки,
не повторяя её раскладку байтов у себя. Цена замерена: 16.96 ГБ весов
превращаются в 18.23 ГБ, то есть развёрнутая таблица стоит 1.27 ГБ.
Это путь совместимости, а не путь скорости: vLLM держит эмбеддинги
упакованными и этой надбавки не платит.
"""
emb = getattr(self, "embed_tokens", None)
if emb is None or getattr(emb, "weight", None) is not None:
return
if not hasattr(emb, "weight_packed"):
return
from compressed_tensors.compressors import BaseCompressor
compressor = BaseCompressor.get_value_from_registry("pack-quantized")
if isinstance(compressor, type):
compressor = compressor()
compressor.decompress_module(emb)
# Снять с модуля признак квантованного. Иначе библиотека увидит его в
# своём списке на первом прямом проходе (`ct_decompress_hook`
# разворачивает всю модель разом) и упадёт с `KeyError: weight_packed`,
# потому что упакованного тензора здесь уже нет. Замерено 26.07.2026.
#
# Именно удалить атрибут, а не обнулить: `is_module_quantized` сперва
# проверяет его наличие, а затем читает поле `weights`, поэтому
# присвоенный None даёт `AttributeError: 'NoneType' object has no
# attribute 'weights'`. Тоже замерено.
if hasattr(emb, "quantization_scheme"):
del emb.quantization_scheme
# ---- проход с оверлеем -------------------------------------------------
def forward(self, input_ids=None, attention_mask=None, position_ids=None,
past_key_values=None, inputs_embeds=None, use_cache=None,
**kwargs):
if inputs_embeds is None:
inputs_embeds = self.embed_tokens(input_ids)
text_position_ids = position_ids
if text_position_ids is not None and text_position_ids.ndim == 3:
text_position_ids = text_position_ids[0]
if text_position_ids is None:
seen = past_key_values.get_seq_length() if past_key_values is not None else 0
text_position_ids = torch.arange(
seen, seen + inputs_embeds.shape[1], device=inputs_embeds.device
).unsqueeze(0)
mask_kwargs = {"config": self.config, "input_embeds": inputs_embeds,
"attention_mask": attention_mask,
"past_key_values": past_key_values,
"position_ids": text_position_ids}
try:
masks = {"full_attention": create_causal_mask(**mask_kwargs),
"linear_attention": create_recurrent_attention_mask(**mask_kwargs)}
except TypeError: # старая сигнатура: inputs_embeds вместо input_embeds
mask_kwargs["inputs_embeds"] = mask_kwargs.pop("input_embeds")
masks = {"full_attention": create_causal_mask(**mask_kwargs),
"linear_attention": create_recurrent_attention_mask(**mask_kwargs)}
hidden_states = inputs_embeds
position_embeddings = self.rotary_emb(hidden_states, position_ids)
replay_cache = self._replay_cache(past_key_values)
for i, layer in enumerate(self.layers[: self.config.num_hidden_layers]):
hidden_states = layer(
hidden_states, position_embeddings=position_embeddings,
attention_mask=masks[self.config.layer_types[i]],
position_ids=text_position_ids, past_key_values=past_key_values,
use_cache=use_cache, **kwargs)
name = ("sb_a" if i == SUPERBLOCK_A_ATTENTION
else "sb_b" if i == SUPERBLOCK_B_ATTENTION else None)
if name is not None:
hidden_states = self.altay_overlay.run(
name, hidden_states,
recurrent_mask=masks["linear_attention"],
attention_mask=masks["full_attention"],
position_embeddings=position_embeddings,
position_ids=text_position_ids,
replay_cache=replay_cache, **kwargs)
hidden_states = self.norm(hidden_states)
return Qwen3_5ModelOutputWithPast(last_hidden_state=hidden_states,
past_key_values=past_key_values)
def _replay_cache(self, past_key_values):
"""Отдельный кэш для повторов: веса общие, рекуррентные состояния — нет."""
if past_key_values is None:
return None
cache = getattr(self, "_altay_replay_cache", None)
if cache is None or type(cache) is not type(past_key_values):
cache = type(past_key_values)(config=self.config)
object.__setattr__(self, "_altay_replay_cache", cache)
return cache
class LomonosovZenitAltayModel(Qwen3_5Model):
"""Мультимодальная обёртка: зрение стоковое, текст с оверлеем."""
def __init__(self, config):
super().__init__(config)
self.language_model = LomonosovZenitAltayTextModel(config.text_config)
self.post_init()
class LomonosovZenitAltayForConditionalGeneration(Qwen3_5ForConditionalGeneration):
"""Точка входа для `from_pretrained`.
Класс мультимодальный, и это существенно: чекпоинт держит веса под
`model.language_model.*`, а `AutoModelForCausalLM` строит текстовое дерево
под `model.*`. Замерено, что при этом совпадает ноль ключей из 853 — мимо
идут даже нормировки, которых квантование не касается. Правильный автокласс
здесь `AutoModelForImageTextToText`; он и прописан в `auto_map`.
"""
def __init__(self, config):
super().__init__(config)
self.model = LomonosovZenitAltayModel(config)
self.post_init()
@classmethod
def from_pretrained(cls, *args, **kwargs):
# Распаковка эмбеддингов возможна только после того, как веса легли на
# место: до загрузки распаковывать нечего.
model = super().from_pretrained(*args, **kwargs)
text = getattr(getattr(model, "model", None), "language_model", None)
if hasattr(text, "materialize_packed_embedding"):
text.materialize_packed_embedding()
return model
__all__ = [
"LomonosovZenitAltayForConditionalGeneration",
"LomonosovZenitAltayModel",
"LomonosovZenitAltayTextModel",
"AltayOverlay",
]