Архитектура модели
GigaChat-20B-A3B состоит из следующих деталей:
- Fine-grained Experts + Shared Experts
- Grouped Query Attention
- Rotary Position Embeddings
- RMSNorm
- SwiGLU в MLP
Важно то, что в реализации MoE некоторые эксперты вызываются в зависимости от контекста, а другие используются всегда.
Бенчмарки
Общие английские метрики. Для замера использовался популярный открытый репозиторий LM Evaluation Harness.
Table with columns: Bench, T-lite-0.1(llama 3.0 8B based), Llama-3.1-8B, GigaChat-20B-A3B-base, Gemma-9B| Bench | T-lite-0.1(llama 3.0 8B based) | Llama-3.1-8B | GigaChat-20B-A3B-base | Gemma-9B |
|---|
| MMLU (5-shot) | 62.56 | 65.21 | 63.02 | 70.6 |
| MMLU-pro (5-shot) | 32.19 | 35.7 | 31.41 | 42.85 |
| MMLU-ru (5-shot) | 55.51 | 54.1 | 58.38 | 62.57 |
| BBH (3-shot) | 62.36 | 62.79 | 53.54 | 70.48 |
| ARC-C (25-shot) | 58.19 |
Requirements
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM, GenerationConfig
model_name = "ai-sage/GigaChat-20B-A3B-base"
tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(model_name, trust_remote_code=True, torch_dtype=torch.bfloat16, device_map="auto")
model.generation_config = GenerationConfig.from_pretrained(model_name)
messages = (
"Ниже я написал подробное доказательство теоремы о неподвижной точке:"
)
input_tensor = tokenizer(messages, return_tensors="pt").input_ids
outputs = model.generate(input_tensor.to(model.device))
result = tokenizer.decode(outputs[0][input_tensor.shape[1]:], skip_special_tokens=False)
print(result)
Пример использования через vLLM
from transformers import AutoTokenizer
from vllm import LLM, SamplingParams
model_name = "ai-sage/GigaChat-20B-A3B-base"
llm = LLM(model=model_name, tokenizer=model_name, trust_remote_code=True)
sampling_params = SamplingParams(
temperature=0.3,
max_tokens=8192,
stop_token_ids=[tokenizer.eos_token_id]
)
messages = (
"Ниже я написал подробное доказательство теоремы о неподвижной точке:"
)
outputs = llm.generate(messages, sampling_params=sampling_params)
generated_text = [output.outputs[0].text for output in outputs]
print(generated_text)
Скорость генерации
Table with columns: Model, Total params (B), Active params (B), Req/s, Output Token/s, Total Token/s| Model | Total params (B) | Active params (B) | Req/s | Output Token/s | Total Token/s |
|---|
| Qwen/Qwen1.5-MoE-A2.7B-Chat | 14 | 2,7 | 0,62 | 156,43 | 291,17 |
| deepseek-ai/deepseek-moe-16b-chat | 16 | 2,8 | 0,59 | 149,53 | 285,39 |
| GigaChat-20B-A3B | 20 |