diff --git a/README.md b/README.md index 7c87885..43f00b5 100644 --- a/README.md +++ b/README.md @@ -138,13 +138,12 @@ source activate.fish --fast эксперементаль флаг который ускоряет foreword примерно в x2-x5 ``` -### update - проверка обновления - ### info ```bash Xenith info models/my.xnh # конфигурация и словарь ``` +### update - проверка обновления `gradcheck` стоит запускать после любой правки в `model.cpp` / `backward.cpp` — он ловит ошибку в обратном проходе за секунды. diff --git a/_xenith b/_xenith new file mode 100644 index 0000000..b35779c --- /dev/null +++ b/_xenith @@ -0,0 +1,80 @@ +# _xenith - Zsh completion function for Xenith +_xenith() { + local curcontext="$curcontext" state line + typeset -A opt_args + local commands=( + 'new:создать новую модель из корпуса' + 'train:обучить существующую модель' + 'gen:сгенерировать текст' + 'info:показать конфигурацию модели' + 'gradcheck:проверить градиенты численно' + 'bench:замерить скорость форварда и шага обучения' + 'update:автоматическое обновление с гита' + 'tui:графическо-консольный режим (тест)' + ) + _arguments -C \ + '1: :->command' \ + '*:: :->args' && return 0 + case $state in + command) + _describe -t commands 'команды Xenith' commands + ;; + args) + case $line[1] in + new) + _arguments \ + '--corpus[текст для построения словаря]:файл:_files' \ + '--out[файл модели]:файл:_files -g "*.xnh"' \ + '--vocab[размер словаря]' \ + '--embd[размер эмбеддинга]' \ + '--layers[число слоёв]' \ + '--heads[число голов внимания]' \ + '--ctx[максимальный контекст]' \ + '--ffn[размер FFN]' \ + '--rope[база RoPE]' \ + '--seed[зерно инициализации]' \ + '--untied[использовать отдельную матрицу выхода]' \ + '--view-vocab[вывести словарь нейросети]' + ;; + train) + _arguments \ + '--corpus[обучающий текст]:файл:_files' \ + '--out[куда сохранить]:файл:_files -g "*.xnh"' \ + '--steps[количество шагов обучения]' \ + '--batch[размер батча]' \ + '--block[длина окна контекста]' \ + '--lr[скорость обучения]' \ + '--wd[weight decay]' \ + '--clip[клиппинг нормы градиента]' \ + '--warmup[шаги прогрева]' \ + '--seed[зерно генератора]' \ + '--ckpt-every[сохранять чекпоинт каждые N шагов]' \ + '--val-every[валидация каждые N шагов]' \ + '--val-tokens[размер валидационной выборки]' \ + '--threads[количество потоков (0 = все ядра)]' \ + '--resume[продолжить обучение с последнего чекпоинта]' \ + '--fast[экспериментальный флаг ускорения forward pass]' + ;; + gen) + _arguments \ + '--prompt[стартовый текст]' \ + '--n[количество генерируемых токенов]' \ + '--temp[температура выборки (0 = жадный)]' \ + '--top-k[top-k sampling]' \ + '--top-p[top-p nucleus sampling]' \ + '--seed[зерно генератора]' \ + '--no-stream[не печатать токены по мере генерации]' \ + '--batch[несколько промптов через ;]' \ + '--show-tokens[показывать ID токенов]' + ;; + info|gradcheck|bench|update|tui) + _arguments \ + '--fast[экспериментальный флаг ускорения]' + ;; + esac + ;; + esac +} +if [[ "$funcstack[1]" == "_xenith" ]]; then + _xenith "$@" +fi \ No newline at end of file diff --git a/activate.zsh b/activate.zsh index 1f90df8..2158703 100644 --- a/activate.zsh +++ b/activate.zsh @@ -1,13 +1,13 @@ +# activate.zsh export VIRTUAL_ENV_PROMPT="%F{red}Xenith%f" alias Xenith='./bin/Xenith' if [ -z "${_OLD_VIRTUAL_PS1+x}" ]; then _OLD_VIRTUAL_PS1="$PS1" fi -PS1="%{($VIRTUAL_ENV_PROMPT)%} $_OLD_VIRTUAL_PS1" +PS1="(${VIRTUAL_ENV_PROMPT}) ${_OLD_VIRTUAL_PS1}" deactivate_custom() { unalias Xenith 2>/dev/null unalias deactivate 2>/dev/null - if [ -n "${_OLD_VIRTUAL_PS1+x}" ]; then PS1="$_OLD_VIRTUAL_PS1" unset _OLD_VIRTUAL_PS1 @@ -18,5 +18,10 @@ deactivate_custom() { builtin deactivate 2>/dev/null || true fi unset -f deactivate_custom + compdel Xenith 2>/dev/null + compdel ./bin/Xenith 2>/dev/null } alias deactivate='deactivate_custom' +source "${0:A:h}/_xenith" +compdef _xenith Xenith +compdef _xenith ./bin/Xenith \ No newline at end of file diff --git a/bin/Xenith b/bin/Xenith index 7b0d252..15627c1 100755 Binary files a/bin/Xenith and b/bin/Xenith differ diff --git a/data.txt b/data.txt new file mode 100644 index 0000000..8e1eb3f --- /dev/null +++ b/data.txt @@ -0,0 +1,2 @@ +Привет +Привет \ No newline at end of file diff --git a/main.cpp b/main.cpp index 7c876f2..b845ca8 100644 --- a/main.cpp +++ b/main.cpp @@ -302,15 +302,6 @@ static bool available_new_version() { Version remote_ver = parse_version(remote_str); Version local_ver = parse_version(local_str); - - if (is_newer(remote_ver, local_ver)) - { - std::cout << "Локальная версия: " << local_ver.major << "." - << local_ver.minor << "." << local_ver.patch << std::endl; - std::cout << "Удалённая версия: " << remote_ver.major << "." - << remote_ver.minor << "." << remote_ver.patch << std::endl; - } - return is_newer(remote_ver, local_ver); } @@ -412,7 +403,7 @@ int main(int argc, char** argv) { std::cerr << "Ошибка чтения файла новой версии (Пожалуйста напишите о этом баге)" << std::endl; return 1; } - std::cout << "\033[32mДоступна новая версия \033[33m" << remote_str << "\033[32m\nВыполните \"Xenith update\" чтобы обновится до последней версии\033[39m\n" << std::endl; + std::cout << "\033[32mДоступна новая версия \033[33m" << remote_str << "\033[32m\nВыполните \"Xenith update\" чтобы обновится до последней версии\033[39m" << std::endl; } if (argc < 2) { printf(HELP_TEXT); return 1; } diff --git a/xenith/core/backward.cpp b/xenith/core/backward.cpp index 723f3a8..153afbd 100644 --- a/xenith/core/backward.cpp +++ b/xenith/core/backward.cpp @@ -17,7 +17,7 @@ static float softmax_ce(const float* logits, int V, int target, float* dlogits) } float backward(Params& p, const ForwardCache& fc, - const int* x, const int* y, int n_pairs) { + const int* x, const int* y, int n_pairs, bool fast) { const Config& c = p.cfg; const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn(); const int V = c.vocab_out(); @@ -40,14 +40,14 @@ float backward(Params& p, const ForwardCache& fc, Tensor dxf(B, d); if (p.tied()) { gemm_nn(dlogits, p.wte, dxf); - gemm_tn(dlogits, fc.xf, p.gwte, true); + gemm_tn(dlogits, fc.xf, p.gwte, true, fast); } else { gemm_nn(dlogits, p.lm_head, dxf); - gemm_tn(dlogits, fc.xf, p.glm_head, true); + gemm_tn(dlogits, fc.xf, p.glm_head, true, fast); } Tensor dx; - rmsnorm_backward(fc.x, p.rms_final, fc.inv_rms, dxf, dx, p.grms_final); + rmsnorm_backward(fc.x, p.rms_final, fc.inv_rms, dxf, dx, p.grms_final, fast); Tensor tmp, dxb, dx2b, dha, dh1, dh3, dq, dk, dv, datt, dattout; @@ -56,7 +56,7 @@ float backward(Params& p, const ForwardCache& fc, const Tensor dres = dx; - gemm_tn(lc.ha, dres, p.g2[L], true); // dW2 + gemm_tn(lc.ha, dres, p.g2[L], true, fast); // dW2 gemm_nt(dres, p.w2[L], dha); // dha = dres @ w2^T dh1.resize(B, f); @@ -73,8 +73,8 @@ float backward(Params& p, const ForwardCache& fc, d3[j] = ga * silu(g); } } - gemm_tn(lc.x2b, dh1, p.g1[L], true); // dW1 - gemm_tn(lc.x2b, dh3, p.g3[L], true); // dW3 + gemm_tn(lc.x2b, dh1, p.g1[L], true, fast); // dW1 + gemm_tn(lc.x2b, dh3, p.g3[L], true, fast); // dW3 dx2b.resize(B, d); gemm_nt(dh1, p.w1[L], dx2b); @@ -83,11 +83,11 @@ float backward(Params& p, const ForwardCache& fc, { Tensor g; - rmsnorm_backward(lc.x2, p.rms_ffn[L], lc.inv_rms2, dx2b, g, p.grms_ffn[L]); + rmsnorm_backward(lc.x2, p.rms_ffn[L], lc.inv_rms2, dx2b, g, p.grms_ffn[L], fast); for (size_t i = 0; i < dx.n(); ++i) dx.at(i) += g.at(i); } - gemm_tn(lc.attout, dx, p.go[L], true); + gemm_tn(lc.attout, dx, p.go[L], true, fast); dattout.resize(B, d); gemm_nt(dx, p.wo[L], dattout); @@ -161,9 +161,9 @@ float backward(Params& p, const ForwardCache& fc, } } - gemm_tn(lc.xb, dq, p.gq[L], true); - gemm_tn(lc.xb, dk, p.gk[L], true); - gemm_tn(lc.xb, dv, p.gv[L], true); + gemm_tn(lc.xb, dq, p.gq[L], true, fast); + gemm_tn(lc.xb, dk, p.gk[L], true, fast); + gemm_tn(lc.xb, dv, p.gv[L], true, fast); dxb.resize(B, d); gemm_nt(dq, p.wq[L], dxb); @@ -174,7 +174,7 @@ float backward(Params& p, const ForwardCache& fc, { Tensor g; - rmsnorm_backward(lc.x, p.rms_attn[L], lc.inv_rms, dxb, g, p.grms_attn[L]); + rmsnorm_backward(lc.x, p.rms_attn[L], lc.inv_rms, dxb, g, p.grms_attn[L], fast); for (size_t i = 0; i < dx.n(); ++i) dx.at(i) += g.at(i); } } diff --git a/xenith/core/generate.cpp b/xenith/core/generate.cpp index 097afd5..e1c09a8 100644 --- a/xenith/core/generate.cpp +++ b/xenith/core/generate.cpp @@ -58,7 +58,7 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos, Tensor ln, lnf; for (int L = 0; L < c.n_layer; ++L) { - rmsnorm_forward(x, p.rms_attn[L], c.rms_eps, xb, ln); + rmsnorm_forward(x, p.rms_attn[L], c.rms_eps, xb, ln, fast); gemm_nn(xb, p.wq[L], q); gemm_nn(xb, p.wk[L], k); gemm_nn(xb, p.wv[L], v); @@ -99,7 +99,7 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos, gemm_nn(attout, p.wo[L], proj); for (int j = 0; j < d; ++j) x.at(j) += proj.at(j); - rmsnorm_forward(x, p.rms_ffn[L], c.rms_eps, x2b, ln); + rmsnorm_forward(x, p.rms_ffn[L], c.rms_eps, x2b, ln, fast); gemm_nn(x2b, p.w1[L], h1); gemm_nn(x2b, p.w3[L], h3); if (fast) attout.resize_fast(1, f); @@ -109,7 +109,7 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos, for (int j = 0; j < d; ++j) x.at(j) += fo.at(j); } - rmsnorm_forward(x, p.rms_final, c.rms_eps, lnf, ln); + rmsnorm_forward(x, p.rms_final, c.rms_eps, lnf, ln, fast); if (fast) attout.resize_fast(1, c.vocab_out()); else attout.resize(1, c.vocab_out()); if (p.tied()) gemm_nt(lnf, p.wte, logits); diff --git a/xenith/core/gradcheck.cpp b/xenith/core/gradcheck.cpp index 1220c60..2c3fa2c 100644 --- a/xenith/core/gradcheck.cpp +++ b/xenith/core/gradcheck.cpp @@ -54,7 +54,7 @@ int run_gradcheck(bool verbose, bool fast) { ForwardCache fc; forward(p, x.data(), n, n, fc, fast); - backward(p, fc, x.data(), y.data(), n); + backward(p, fc, x.data(), y.data(), n, fast); auto tensors = p.all_w(); auto grads = p.all_grads(); @@ -130,7 +130,7 @@ int run_bench(Checkpoint& ck, int iters, int block, bool fast) { // прогрев ForwardCache fc; forward(ck.params, x.data(), n, n, fc, fast); - backward(ck.params, fc, x.data(), y.data(), n); + backward(ck.params, fc, x.data(), y.data(), n, fast); auto t0 = std::chrono::steady_clock::now(); for (int i = 0; i < iters; ++i) forward(ck.params, x.data(), n, n, fc, fast); @@ -141,7 +141,7 @@ int run_bench(Checkpoint& ck, int iters, int block, bool fast) { t0 = std::chrono::steady_clock::now(); for (int i = 0; i < iters; ++i) { forward(ck.params, x.data(), n, n, fc, fast); - backward(ck.params, fc, x.data(), y.data(), n); + backward(ck.params, fc, x.data(), y.data(), n, fast); } t1 = std::chrono::steady_clock::now(); double d_all = std::chrono::duration(t1 - t0).count(); diff --git a/xenith/core/model.cpp b/xenith/core/model.cpp index 3365098..7a4ed9a 100644 --- a/xenith/core/model.cpp +++ b/xenith/core/model.cpp @@ -240,12 +240,46 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward for (int L = 0; L < c.n_layer; ++L) { LayerCache& lc = fc.layers[L]; - if (fast) fc.x.resize_fast(B, d); - else fc.x.resize(B, d); + if (fast) + { + lc.x.resize_fast(B, d); + lc.xb.resize_fast(B, d); + lc.q.resize_fast(B, d); + lc.k.resize_fast(B, d); + lc.v.resize_fast(B, d); + lc.attout.resize_fast(B, d); + lc.proj.resize_fast(B, d); - for (size_t i = 0; i < fc.x.n(); ++i) lc.x.at(i) = fc.x.at(i); + lc.x2.resize_fast(B, d); + lc.x2b.resize_fast(B, d); + lc.h1.resize_fast(B, f); + lc.h3.resize_fast(B, f); + lc.ha.resize_fast(B, f); + lc.fo.resize_fast(B, d); + } else + { + lc.x.resize(B, d); + lc.xb.resize(B, d); + lc.q.resize(B, d); + lc.k.resize(B, d); + lc.v.resize(B, d); + lc.attout.resize(B, d); + lc.proj.resize(B, d); - rmsnorm_forward(lc.x, p.rms_attn[L], c.rms_eps, lc.xb, lc.inv_rms); + lc.x2.resize(B, d); + lc.x2b.resize(B, d); + lc.h1.resize(B, f); + lc.h3.resize(B, f); + lc.ha.resize(B, f); + lc.fo.resize(B, d); + } + + lc.att.resize3d(n_seq * H, S, S); + + if (fast) fc.xf.resize_fast(B, d); + else fc.xf.resize(B, d); + + rmsnorm_forward(lc.x, p.rms_attn[L], c.rms_eps, lc.xb, lc.inv_rms, fast); gemm_nn(lc.xb, p.wq[L], lc.q); gemm_nn(lc.xb, p.wk[L], lc.k); gemm_nn(lc.xb, p.wv[L], lc.v); @@ -316,7 +350,7 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward else fc.x.resize(B, d); for (size_t i = 0; i < fc.x.n(); ++i) lc.x2.at(i) = fc.x.at(i); - rmsnorm_forward(lc.x2, p.rms_ffn[L], c.rms_eps, lc.x2b, lc.inv_rms2); + rmsnorm_forward(lc.x2, p.rms_ffn[L], c.rms_eps, lc.x2b, lc.inv_rms2, fast); gemm_nn(lc.x2b, p.w1[L], lc.h1); gemm_nn(lc.x2b, p.w3[L], lc.h3); if (fast) fc.x.resize_fast(B, f); @@ -327,7 +361,7 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward for (size_t i = 0; i < fc.x.n(); ++i) fc.x.at(i) += lc.fo.at(i); } - rmsnorm_forward(fc.x, p.rms_final, c.rms_eps, fc.xf, fc.inv_rms); + rmsnorm_forward(fc.x, p.rms_final, c.rms_eps, fc.xf, fc.inv_rms, fast); if (fast) fc.x.resize_fast(B, c.vocab_out()); else fc.x.resize(B, c.vocab_out()); if (p.tied()) gemm_nt(fc.xf, p.wte, fc.logits); diff --git a/xenith/core/tensor.hpp b/xenith/core/tensor.hpp index 4fb7f2e..947488c 100644 --- a/xenith/core/tensor.hpp +++ b/xenith/core/tensor.hpp @@ -132,10 +132,11 @@ inline void gemm_nt(const Tensor& A, const Tensor& B, Tensor& C) { }); } -inline void gemm_tn(const Tensor& A, const Tensor& B, Tensor& C, bool accumulate = false) { +inline void gemm_tn(const Tensor& A, const Tensor& B, Tensor& C, bool accumulate, bool fast) { const int N = A.R(), K = A.C(), M = B.C(); if (!accumulate) { - C.resize(K, M); + if (fast) C.resize_fast(K, M); + else C.resize(K, M); } else { if (C.R() != K || C.C() != M) { std::fprintf(stderr, @@ -167,10 +168,12 @@ inline float dsilu(float x) { inline void rmsnorm_forward(const Tensor& x, const Tensor& w, float eps, - Tensor& y, Tensor& inv_rms) { + Tensor& y, Tensor& inv_rms, bool fast) { const int N = x.R(), D = x.C(); - y.resize(N, D); - inv_rms.resize(N, 1); + if (fast) y.resize_fast(N, D); + else y.resize(N, D); + if (fast) inv_rms.resize_fast(N, 1); + else inv_rms.resize(N, 1); for (int n = 0; n < N; ++n) { const float* xr = x.row(n); float ss = 0.0f; @@ -183,10 +186,12 @@ inline void rmsnorm_forward(const Tensor& x, const Tensor& w, float eps, } inline void rmsnorm_backward(const Tensor& x, const Tensor& w, const Tensor& inv_rms, - const Tensor& dy, Tensor& dx, Tensor& dw) { + const Tensor& dy, Tensor& dx, Tensor& dw, bool fast) { const int N = x.R(), D = x.C(); - dx.resize(N, D); - dw.resize(1, D); + if (fast) dx.resize_fast(N, D); + else dx.resize(N, D); + if (fast) dw.resize_fast(1, D); + else dw.resize(1, D); dw.zero(); for (int n = 0; n < N; ++n) { const float* xr = x.row(n); diff --git a/xenith/core/train.hpp b/xenith/core/train.hpp index db714b6..d82c729 100644 --- a/xenith/core/train.hpp +++ b/xenith/core/train.hpp @@ -6,7 +6,7 @@ namespace xt { float backward(Params& p, const ForwardCache& fc, - const int* x, const int* y, int n_pairs); + const int* x, const int* y, int n_pairs, bool fast); float softmax_eval(const Params& p, const ForwardCache& fc, const int* x, const int* y, int n_pairs, Tensor& per_token_loss); diff --git a/xenith/core/trainer.cpp b/xenith/core/trainer.cpp index ede0c2c..132dd7f 100644 --- a/xenith/core/trainer.cpp +++ b/xenith/core/trainer.cpp @@ -80,7 +80,7 @@ TrainStats train_model(Params& p, const Tokenizer& tok, const std::string& corpu } forward(p, xb.data(), B * block, block, fc, tc.fast); - float loss = backward(p, fc, xb.data(), yb.data(), B * block); + float loss = backward(p, fc, xb.data(), yb.data(), B * block, tc.fast); p.adam_step(lr_at(tc, step), tc.beta1, tc.beta2, tc.eps, tc.weight_decay, step + 1, tc.clip);