Xenith/xenith/core/model.h
2026-09-29 19:55:06 +07:00

105 lines
4.4 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// model.h — конфигурация, параметры и прямой проход трансформера
#pragma once
#include "tensor.h"
#include <string>
#include <vector>
namespace xt {
// Магия формата: 'X','N','H','1' в little-endian
constexpr uint32_t XNH_MAGIC = 0x31484E58;
struct Config {
int vocab_size = 256;
int n_layer = 4;
int n_head = 4;
int n_embd = 128;
int block_size = 128; // максимальный контекст
int ffn_dim = 0; // 0 => 4*n_embd
int rope_base = 10000;
float rms_eps = 1e-5f;
float init_std = 0.02f;
int n_vocab_out = 0; // 0 => tied (используем wte как выходную матрицу)
int head_dim() const { return n_embd / n_head; }
int ffn() const { return ffn_dim > 0 ? ffn_dim : 4 * n_embd; }
int vocab_out() const { return n_vocab_out > 0 ? n_vocab_out : vocab_size; }
bool valid(std::string& err) const;
};
// Все обучаемые веса. Именно этот список сериализуется в файл модели.
struct Params {
Config cfg;
Tensor wte; // [vocab, n_embd] — вход
Tensor lm_head; // [n_embd, vocab_out] — только если НЕ tied
std::vector<Tensor> wq, wk, wv, wo;
std::vector<Tensor> rms_attn, rms_ffn;
std::vector<Tensor> w1, w2, w3; // SwiGLU
Tensor rms_final; // [1, n_embd]
// Градиенты (тет же порядок, что у all())
Tensor gwte, glm_head, grms_final;
std::vector<Tensor> gq, gk, gv, go, grms_attn, grms_ffn, g1, g2, g3;
// Моменты Adam
std::vector<Tensor> m, v;
size_t n_params() const;
bool tied() const;
void alloc_grads();
void init(uint64_t seed = 1234);
void zero_grad();
void adam_step(float lr, float b1, float b2, float eps, float wd,
int step, float clip);
std::vector<const Tensor*> all() const; // только для чтения/сохранения
std::vector<Tensor*> all_w(); // для изменения весов
std::vector<Tensor*> all_grads();
};
// Кэш прямого прохода, нужен обратному
struct LayerCache {
Tensor x; // вход в слой
Tensor xb; // нормализованный вход (attention)
Tensor q, k, v; // [B, n_embd], RoPE применён к q,k
Tensor att; // [n_head, B, B] — веса внимания
Tensor attout; // [B, n_embd] — конкатенация голов
Tensor proj; // [B, n_embd] — после wo
Tensor x2; // после residual attention
Tensor x2b; // нормализация перед FFN
Tensor h1, h3, ha;
Tensor fo; // [B, n_embd] — после w2
Tensor inv_rms, inv_rms2;
};
struct ForwardCache {
std::vector<LayerCache> layers;
Tensor x; // текущий residual-поток
Tensor xf; // после финального RMSNorm
Tensor logits; // [B, vocab_out]
Tensor inv_rms; // финальная нормализация
std::vector<float> rope_cos, rope_sin;
int block = 0; // длина одной последовательности
int seq_len = 0; // то же, явно
int n_seq = 1; // сколько последовательностей в батче
};
void rope_tables(int block, int head_dim, int base,
std::vector<float>& cs, std::vector<float>& sn);
// Полный прямой проход.
// n_tok — сколько токенов всего (может быть кратно block_size: батч)
// seq_len — длина ОДНОЙ последовательности (<= block_size).
// Внимание и RoPE строятся независимо для каждой
// последовательности, токены разных батчей не видят друг друга.
void forward(const Params& p, const int* tokens, int n_tok, int seq_len, ForwardCache& fc);
// Обёртка для одиночной последовательности
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc);
// Только последний ряд логитов
void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits);
} // namespace xt