Big-CoMAttn
Небольшая-предтрен языковая модель на русском языке, обученная с нуля на эксперементальной - гибридной архитектуре Mamba + Attention.
Архитектура
Гибридная архитектура,Mamba (State Space Model)aceGrouped Query Attention (GQA)tion (GQA)**. Общая схема:
┌────────────────────────────────────────────────────┐
│ Layer 1: MambaBlock │
│ LayerNorm → Mamba │
├────────────────────────────────────────────────────┤
│ Layer 2: TransformerBlockBig │
│ LayerNorm → GQA Attention │
│ LayerNorm → MLP (6×) │
├────────────────────────────────────────────────────┤
│ Layer 3: MambaBlock │
│ LayerNorm → Mamba │
├────────────────────────────────────────────────────┤
│ Layer 4: MambaBlock │
│ LayerNorm → Mamba │
├────────────────────────────────────────────────────┤
│ Layer 5: AttnMambaBlock │
│ LayerNorm → GQA Attention │
│ LayerNorm → Mamba │
├────────────────────────────────────────────────────┤
│ Layer 6: MambaBlock │
│ LayerNorm → Mamba │
├────────────────────────────────────────────────────┤
│ Layer 7: MambaBlock │
│ LayerNorm → Mamba │
├────────────────────────────────────────────────────┤
│ Layer 8: FinalBlock │
│ LayerNorm → GQA Attention │
│ LayerNorm → MLP (2×) │
└────────────────────────────────────────────────────┘
lm_head.weight толстенький пирожок 31% от всей модели что было ошибкой, в последствии стоит использовать свой токенизатор. tying - многократно ухудшал результаты было принято решение отказаться.
Mamba
Все Mamba-слои используют одинаковые параметры:
Mamba( d_model=512, d_state=16, d_conv=4, expand=2, ) Grouped Query Attention (GQA)
Все attention-слои используют GQA с отвязанным head_dim:
GQAAttention( d_model=512, n_heads_q=16, n_heads_kv=4, # GQA ratio = 4:1 head_dim=64, # отвязано от d_model // n_heads_q dropout=0.1, )
Используется F.scaled_dot_product_attention с broadcasting для GQA (без repeat_interleave — без копий K/V).
MLP
· Слой 2: Linear(512 → 3072) → GELU → Linear(3072 → 512) — 6× · Слой 8: Linear(512 → 1024) → GELU → Linear(1024 → 512) — 2×
Узкий MLP в финальном слое выбран намеренно(эмпирически): на выходе модели нужна только нелинейность.
Инициализация
GPT-стиль инициализации (N(0, 0.02)) - Attention.
Пример использования
import os
os.environ["TOKENIZERS_PARALLELISM"] = "false"
import torch
import torch.nn.functional as F
from huggingface_hub import snapshot_download
from safetensors.torch import load_file
from transformers import AutoTokenizer
from mambaMain import build_model
REPO_ID = "algorithms-learning/Big-CoMAttn"
CACHE_DIR = "./hf_cache"
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
TEMPERATURE = 0.3
TOP_K = 50
TOP_P = 0.95
REPETITION_PENALTY = 1.2
MAX_NEW_TOKENS = 150
def download_model(repo_id=REPO_ID, cache_dir=CACHE_DIR):
path = snapshot_download(
repo_id=repo_id,
cache_dir=cache_dir,
allow_patterns=[
"*.safetensors",
"*.json",
"*.model",
"README.md",
],
)
return path
def load_model_and_tokenizer(model_dir):
tokenizer = AutoTokenizer.from_pretrained(model_dir)
tokenizer.pad_token = tokenizer.eos_token
print(f"Tokenizer loaded. Vocab: {len(tokenizer)}")
model = build_model(len(tokenizer), config={
"d_model": 512,
"max_len": 1024,
"n_heads_q": 16,
"n_heads_kv": 4,
"head_dim": 64,
"dropout": 0.1,
}).to(DEVICE)
state_dict = load_file(
os.path.join(model_dir, "model.safetensors"),
device=DEVICE,
)
model.load_state_dict(state_dict)
model.eval()
n_params = sum(p.numel() for p in model.parameters())
print(f"Model loaded. Params: {n_params/1e6:.2f}M")
return model, tokenizer
def generate(model, tokenizer, prompt,
max_new_tokens=MAX_NEW_TOKENS,
temperature=TEMPERATURE,
top_k=TOP_K, top_p=TOP_P,
rep_penalty=REPETITION_PENALTY):
"""Sampling: temperature + top-k + top-p + repetition penalty."""
ids = tokenizer.encode(prompt, return_tensors="pt").to(DEVICE)
with torch.no_grad():
for _ in range(max_new_tokens):
input_ids = ids[:, -1024:]
logits = model(input_ids)[:, -1, :].float()
if rep_penalty != 1.0:
for token_id in set(ids[0, -50:].tolist()):
if logits[0, token_id] > 0:
logits[0, token_id] /= rep_penalty
else:
logits[0, token_id] *= rep_penalty
logits = logits / temperature
if top_k > 0:
top_k_vals, _ = torch.topk(logits, top_k)
min_val = top_k_vals[0, -1]
logits[logits < min_val] = float("-inf")
if top_p < 1.0:
sorted_logits, sorted_idx = torch.sort(logits, descending=True)
probs = F.softmax(sorted_logits, dim=-1)
cum_probs = torch.cumsum(probs, dim=-1)
sorted_mask = cum_probs > top_p
sorted_mask[..., 1:] = sorted_mask[..., :-1].clone()
sorted_mask[..., 0] = False
indices_to_remove = sorted_mask.scatter(1, sorted_idx, sorted_mask)
logits[indices_to_remove] = float("-inf")
probs = F.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
ids = torch.cat([ids, next_token], dim=1)
if next_token.item() == tokenizer.eos_token_id:
break
return tokenizer.decode(ids[0], skip_special_tokens=True)
def main():
model_dir = download_model()
print(f"Model dir: {model_dir}\n")
model, tokenizer = load_model_and_tokenizer(model_dir)
print()
prompts = [
"В твоей модели мира число звёзд бесконечно",
"Привет, как дела? Я хотел тебе сказать",
"Он посмотрел на неё и понял, что",
"Вчера в Москве произошло странное событие",
"Анон, объясни мне пожалуйста",
]
for i, p in enumerate(prompts, 1):
print("=" * 60)
print(f"[{i}] PROMPT: {p}")
print("-" * 60)
result = generate(model, tokenizer, p)
print(result)
print()
if __name__ == "__main__":
main()