90 lines
2.2 KiB
C++
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);
|
|
|
|
} |