Team Ai
Modelpublic

algorithms-learning/Big-CoMAttn

sourceHugging Faceapache-2.0updated 4h agoView on Hugging Face
0likes
Model Card

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( dmodel=512, dstate=16, d_conv=4, expand=2, ) Grouped Query Attention (GQA)

Все attention-слои используют GQA с отвязанным head_dim:

GQAAttention( dmodel=512, nheadsq=16, nheadskv=4, # GQA ratio = 4:1 headdim=64, # отвязано от dmodel // nheads_q dropout=0.1, )

Используется F.scaleddotproductattention с broadcasting для GQA (без repeatinterleave — без копий 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.

Пример использования

python
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()