Xenith/xenith/core/model.hpp
2026-10-02 19:30:49 +07:00

90 lines
2.2 KiB
C++

#pragma once
#include "tensor.hpp"
#include <string>
#include <vector>
namespace xt {
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;
int rope_base = 10000;
float rms_eps = 1e-5f;
float init_std = 0.02f;
int n_vocab_out = 0;
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;
Tensor lm_head;
std::vector<Tensor> wq, wk, wv, wo;
std::vector<Tensor> rms_attn, rms_ffn;
std::vector<Tensor> w1, w2, w3;
Tensor rms_final;
Tensor gwte, glm_head, grms_final;
std::vector<Tensor> gq, gk, gv, go, grms_attn, grms_ffn, g1, g2, g3;
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;
Tensor q, k, v;
Tensor att;
Tensor attout;
Tensor proj;
Tensor x2;
Tensor x2b;
Tensor h1, h3, ha;
Tensor fo;
Tensor inv_rms, inv_rms2;
};
struct ForwardCache {
std::vector<LayerCache> layers;
Tensor x;
Tensor xf;
Tensor logits;
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);
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);
}