--fast
This commit is contained in:
parent
ecf33beb01
commit
f0c1d23121
16
README.md
16
README.md
@ -123,14 +123,24 @@ source activate.fish
|
||||
--conf NAME.conf конфиг с настройками и списком моделей
|
||||
```
|
||||
|
||||
### gradcheck - проверить градиенты численно
|
||||
|
||||
```
|
||||
--fast эксперементаль флаг который ускоряет foreword примерно в x2-x5
|
||||
```
|
||||
|
||||
### bench - замерить скорость форварда и шага обучения
|
||||
|
||||
```
|
||||
--fast эксперементаль флаг который ускоряет foreword примерно в x2-x5
|
||||
```
|
||||
|
||||
### update - проверка обновления
|
||||
|
||||
### info / gradcheck / bench
|
||||
### info
|
||||
|
||||
```bash
|
||||
Xenith info models/my.xnh # конфигурация и словарь
|
||||
Xenith gradcheck # градиенты против численных
|
||||
Xenith bench models/my.xnh # скорость
|
||||
```
|
||||
|
||||
`gradcheck` стоит запускать после любой правки в `model.cpp` / `backward.cpp` —
|
||||
|
||||
BIN
bin/Xenith
BIN
bin/Xenith
Binary file not shown.
20
help.h
20
help.h
@ -91,7 +91,8 @@
|
||||
" --val-every N валидация каждые N шагов (0 = нет)\n" \
|
||||
" --val-tokens N размер валидации (по умолчанию 20000)\n" \
|
||||
" --threads N потоков (0 = все ядра)\n" \
|
||||
" --resume продолжить с сохранённого step (Adam-state из файла)\n"
|
||||
" --resume продолжить с сохранённого step (Adam-state из файла)\n" \
|
||||
" --fast эксперементаль флаг который ускоряет foreword примерно в x2-x5\n"
|
||||
#else
|
||||
#define OPT_TRAIN ""
|
||||
#endif
|
||||
@ -109,6 +110,18 @@
|
||||
#else
|
||||
#define OPT_GEN ""
|
||||
#endif
|
||||
#ifdef GRADCHECK
|
||||
#define OPT_GRADCHECK "\nОПЦИИ GRADCHECK\n" \
|
||||
" --fast эксперементаль флаг который ускоряет foreword примерно в x2-x5\n"
|
||||
#else
|
||||
#define OPT_GRADCHECK ""
|
||||
#endif
|
||||
#ifdef BENCH
|
||||
#define OPT_BENCH "\nОПЦИИ BENCH\n" \
|
||||
" --fast эксперементаль флаг который ускоряет foreword примерно в x2-x5\n"
|
||||
#else
|
||||
#define OPT_BENCH ""
|
||||
#endif
|
||||
#ifdef OLLAMA
|
||||
#define OPT_OLLAMA "\nОПЦИИ ollama\n" \
|
||||
" --port PORT порт на котором будет открыт сервер\n" \
|
||||
@ -133,6 +146,9 @@
|
||||
OPT_NEW \
|
||||
OPT_TRAIN \
|
||||
OPT_GEN \
|
||||
OPT_OLLAMA
|
||||
OPT_GRADCHECK \
|
||||
OPT_BENCH \
|
||||
OPT_OLLAMA \
|
||||
"\n"
|
||||
|
||||
#endif
|
||||
59
main.cpp
59
main.cpp
@ -1,7 +1,7 @@
|
||||
#include "xenith/core/checkpoint.hpp"
|
||||
#include "xenith/core/generate.hpp"
|
||||
#include "xenith/core/train.hpp"
|
||||
#include "xenith/core/gradcheck.h"
|
||||
#include "xenith/core/gradcheck.hpp"
|
||||
#include "xenith/ui/tui.hpp"
|
||||
#include <cstdio>
|
||||
#include "help.h"
|
||||
@ -149,6 +149,7 @@ static int cmd_train(const Args& a) {
|
||||
tc.val_every = a.geti("val-every", 0);
|
||||
tc.val_tokens = a.geti("val-tokens", 20000);
|
||||
tc.threads = a.geti("threads", 0);
|
||||
tc.fast = a.has("fast");
|
||||
|
||||
if (tc.threads > 0) set_threads(tc.threads);
|
||||
if (tc.block > ck.params.cfg.block_size) tc.block = ck.params.cfg.block_size;
|
||||
@ -229,7 +230,7 @@ static int cmd_info(const Args& a) {
|
||||
}
|
||||
|
||||
static int cmd_gradcheck(const Args& a) {
|
||||
return run_gradcheck(a.has("verbose"));
|
||||
return run_gradcheck(a.has("verbose"), a.has("fast"));
|
||||
}
|
||||
|
||||
static int cmd_bench(const Args& a) {
|
||||
@ -237,7 +238,7 @@ static int cmd_bench(const Args& a) {
|
||||
Checkpoint ck;
|
||||
std::string err;
|
||||
if (!load_checkpoint(a.pos[0], ck, err)) { fprintf(stderr, "bench: %s\n", err.c_str()); return 1; }
|
||||
return run_bench(ck, a.geti("iters", 20), a.geti("block", 64));
|
||||
return run_bench(ck, a.geti("iters", 20), a.geti("block", 64), a.has("fast"));
|
||||
}
|
||||
|
||||
static int cmd_ollama(const Args& a) {
|
||||
@ -302,10 +303,13 @@ static bool available_new_version() {
|
||||
Version remote_ver = parse_version(remote_str);
|
||||
Version local_ver = parse_version(local_str);
|
||||
|
||||
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;
|
||||
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);
|
||||
}
|
||||
@ -336,7 +340,17 @@ static int cmd_update(const Args& a) {
|
||||
const std::string REPO_URL = "https://git.bipfr.ru/BIPfR/Xenith.git";
|
||||
const std::string TMP_CLONE = "/tmp/xenith_update";
|
||||
|
||||
system(("rm -rf " + TMP_CLONE).c_str());
|
||||
int ret = system(("rm -rf " + TMP_CLONE).c_str());
|
||||
if (ret == -1)
|
||||
{
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
} else {
|
||||
int exit_code = WEXITSTATUS(ret);
|
||||
|
||||
if (exit_code != 0) {
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
}
|
||||
}
|
||||
|
||||
std::cout << "Скачивание новой версии..." << std::endl;
|
||||
int result = system(("git clone --depth 1 " + REPO_URL + " " + TMP_CLONE).c_str());
|
||||
@ -348,19 +362,40 @@ static int cmd_update(const Args& a) {
|
||||
result = system(("cp -ra " + TMP_CLONE + "/. .").c_str());
|
||||
if (result != 0) {
|
||||
std::cerr << "Ошибка копирования файлов" << std::endl;
|
||||
system(("rm -rf " + TMP_CLONE).c_str());
|
||||
ret = system(("rm -rf " + TMP_CLONE).c_str());
|
||||
if (ret == -1)
|
||||
{
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
} else {
|
||||
int exit_code = WEXITSTATUS(ret);
|
||||
|
||||
if (exit_code != 0) {
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
}
|
||||
}
|
||||
return -1;
|
||||
}
|
||||
system(("rm -rf " + TMP_CLONE).c_str());
|
||||
ret = system(("rm -rf " + TMP_CLONE).c_str());
|
||||
if (ret == -1)
|
||||
{
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
} else {
|
||||
int exit_code = WEXITSTATUS(ret);
|
||||
|
||||
if (exit_code != 0) {
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
}
|
||||
}
|
||||
std::cout << "Обновление завершено!" << std::endl;
|
||||
return 0;
|
||||
}
|
||||
|
||||
static int cmd_tui(const Args& a)
|
||||
{
|
||||
while (ui::app.is_running())
|
||||
auto app = ui::Tui();
|
||||
while (app.is_running())
|
||||
{
|
||||
ui::app.update();
|
||||
app.update();
|
||||
}
|
||||
std::cout << "\033[2J\033[H" << std::flush;
|
||||
return 0;
|
||||
|
||||
@ -1 +1 @@
|
||||
0.0.4
|
||||
0.0.5
|
||||
@ -51,7 +51,14 @@ Model load_from_bfr(const std::string& filename)
|
||||
}
|
||||
|
||||
uint32_t map_size;
|
||||
fread(&map_size, sizeof(map_size), 1, mdl.fd);
|
||||
if (fread(&map_size, sizeof(map_size), 1, mdl.fd))
|
||||
{
|
||||
|
||||
} else
|
||||
{
|
||||
std::cerr << "Ошибка чтения файла" << std::endl;
|
||||
return mdl;
|
||||
}
|
||||
|
||||
if (map_size != mdl.config.vocab_size) {
|
||||
std::cerr << "Warning: vocab size mismatch! "
|
||||
|
||||
@ -241,7 +241,7 @@ bool load_checkpoint(const std::string& path, Checkpoint& ck, std::string& err)
|
||||
if (magic != XNH_MAGIC) {
|
||||
char buf[64];
|
||||
std::snprintf(buf, sizeof buf,
|
||||
"плохая магия 0x%08X (ожидалась 0x%08X) — это не Xenith-модель",
|
||||
"Файл не является поддерживаемым",
|
||||
magic, XNH_MAGIC);
|
||||
err = buf;
|
||||
return false;
|
||||
|
||||
@ -44,7 +44,7 @@ void rope_apply(float* x, int hd, int pos, const std::vector<float>& cs,
|
||||
|
||||
void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
||||
const std::vector<float>& cs, const std::vector<float>& sn,
|
||||
Tensor& logits) {
|
||||
Tensor& logits, 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 float scale = 1.0f / std::sqrt((float)hd);
|
||||
@ -71,8 +71,8 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
||||
for (int j = 0; j < hd; ++j) { kp[j] = k.at(h * hd + j); vp[j] = v.at(h * hd + j); }
|
||||
}
|
||||
|
||||
// attention по всем позициям 0..pos
|
||||
attout.resize(1, d);
|
||||
if (fast) attout.resize_fast(1, d);
|
||||
else attout.resize(1, d);
|
||||
attout.zero();
|
||||
for (int h = 0; h < H; ++h) {
|
||||
const float* qr = q.row(0) + h * hd;
|
||||
@ -102,14 +102,16 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
||||
rmsnorm_forward(x, p.rms_ffn[L], c.rms_eps, x2b, ln);
|
||||
gemm_nn(x2b, p.w1[L], h1);
|
||||
gemm_nn(x2b, p.w3[L], h3);
|
||||
ha.resize(1, f);
|
||||
if (fast) attout.resize_fast(1, f);
|
||||
else attout.resize(1, f);
|
||||
for (size_t i = 0; i < ha.n(); ++i) ha.at(i) = silu(h1.at(i)) * h3.at(i);
|
||||
gemm_nn(ha, p.w2[L], fo);
|
||||
for (int j = 0; j < d; ++j) x.at(j) += fo.at(j);
|
||||
}
|
||||
|
||||
rmsnorm_forward(x, p.rms_final, c.rms_eps, lnf, ln);
|
||||
logits.resize(1, c.vocab_out());
|
||||
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);
|
||||
else gemm_nn(lnf, p.lm_head, logits);
|
||||
}
|
||||
@ -184,7 +186,7 @@ int sample(const Tensor& logits, const GenConfig& gc, Rng& rng, const std::vecto
|
||||
|
||||
std::string generate(const Params& p, const Tokenizer& tok,
|
||||
const std::string& prompt, const GenConfig& gc,
|
||||
std::vector<int>* out_ids) {
|
||||
std::vector<int>* out_ids, bool fast) {
|
||||
const Config& c = p.cfg;
|
||||
KvCache kv;
|
||||
kv.init(c.n_layer, c.block_size, c.n_head, c.head_dim());
|
||||
@ -213,13 +215,13 @@ std::string generate(const Params& p, const Tokenizer& tok,
|
||||
ids.resize(keep);
|
||||
kv.reset();
|
||||
for (int t = 0; t < keep; ++t) {
|
||||
forward_token(p, kv, ids[t], kv.len, cs, sn, logits);
|
||||
forward_token(p, kv, ids[t], kv.len, cs, sn, logits, fast);
|
||||
kv.len++;
|
||||
}
|
||||
}
|
||||
|
||||
const int last = ids.back();
|
||||
forward_token(p, kv, last, kv.len, cs, sn, logits);
|
||||
forward_token(p, kv, last, kv.len, cs, sn, logits, fast);
|
||||
kv.len++;
|
||||
|
||||
int next = sample(logits, gc, rng, ids);
|
||||
|
||||
@ -18,7 +18,7 @@ struct GenConfig {
|
||||
|
||||
std::string generate(const Params& p, const Tokenizer& tok,
|
||||
const std::string& prompt, const GenConfig& gc,
|
||||
std::vector<int>* out_ids = nullptr);
|
||||
std::vector<int>* out_ids = nullptr, bool fast = false);
|
||||
|
||||
std::vector<std::string> generate_batch(const Params& p, const Tokenizer& tok,
|
||||
const std::vector<std::string>& prompts,
|
||||
|
||||
@ -1,6 +1,6 @@
|
||||
// gradcheck.cpp — проверяем, что обратный проход совпадает с численным
|
||||
// и заодно меряем скорость.
|
||||
#include "gradcheck.h"
|
||||
#include "gradcheck.hpp"
|
||||
#include "generate.hpp"
|
||||
#include "train.hpp"
|
||||
#include <cstdio>
|
||||
@ -14,9 +14,9 @@ namespace xt {
|
||||
namespace {
|
||||
|
||||
// loss на фиксированном батче, без обратного прохода
|
||||
float eval_loss(const Params& p, const std::vector<int>& x, const std::vector<int>& y, int n) {
|
||||
float eval_loss(const Params& p, const std::vector<int>& x, const std::vector<int>& y, int n, bool fast) {
|
||||
ForwardCache fc;
|
||||
forward(p, x.data(), n, fc);
|
||||
forward(p, x.data(), n, fc, fast);
|
||||
Tensor loss;
|
||||
return softmax_eval(p, fc, x.data(), y.data(), n, loss);
|
||||
}
|
||||
@ -30,7 +30,7 @@ struct Probe {
|
||||
|
||||
}
|
||||
|
||||
int run_gradcheck(bool verbose) {
|
||||
int run_gradcheck(bool verbose, bool fast) {
|
||||
Config c;
|
||||
c.vocab_size = 24;
|
||||
c.n_layer = 2;
|
||||
@ -53,7 +53,7 @@ int run_gradcheck(bool verbose) {
|
||||
}
|
||||
|
||||
ForwardCache fc;
|
||||
forward(p, x.data(), n, n, fc);
|
||||
forward(p, x.data(), n, n, fc, fast);
|
||||
backward(p, fc, x.data(), y.data(), n);
|
||||
|
||||
auto tensors = p.all_w();
|
||||
@ -80,9 +80,9 @@ int run_gradcheck(bool verbose) {
|
||||
const float orig = W.at(i);
|
||||
|
||||
W.at(i) = orig + eps;
|
||||
const float lp = eval_loss(p, x, y, n);
|
||||
const float lp = eval_loss(p, x, y, n, fast);
|
||||
W.at(i) = orig - eps;
|
||||
const float lm = eval_loss(p, x, y, n);
|
||||
const float lm = eval_loss(p, x, y, n, fast);
|
||||
W.at(i) = orig;
|
||||
|
||||
const float num = (lp - lm) / (2.0f * eps);
|
||||
@ -112,7 +112,7 @@ int run_gradcheck(bool verbose) {
|
||||
return 1;
|
||||
}
|
||||
|
||||
int run_bench(Checkpoint& ck, int iters, int block) {
|
||||
int run_bench(Checkpoint& ck, int iters, int block, bool fast) {
|
||||
const Config& c = ck.params.cfg;
|
||||
if (block > c.block_size) block = c.block_size;
|
||||
const int n = block;
|
||||
@ -129,18 +129,18 @@ int run_bench(Checkpoint& ck, int iters, int block) {
|
||||
|
||||
// прогрев
|
||||
ForwardCache fc;
|
||||
forward(ck.params, x.data(), n, n, fc);
|
||||
forward(ck.params, x.data(), n, n, fc, fast);
|
||||
backward(ck.params, fc, x.data(), y.data(), n);
|
||||
|
||||
auto t0 = std::chrono::steady_clock::now();
|
||||
for (int i = 0; i < iters; ++i) forward(ck.params, x.data(), n, n, fc);
|
||||
for (int i = 0; i < iters; ++i) forward(ck.params, x.data(), n, n, fc, fast);
|
||||
auto t1 = std::chrono::steady_clock::now();
|
||||
double d_fwd = std::chrono::duration<double>(t1 - t0).count();
|
||||
float fwd_tps = (double)n * iters / d_fwd;
|
||||
|
||||
t0 = std::chrono::steady_clock::now();
|
||||
for (int i = 0; i < iters; ++i) {
|
||||
forward(ck.params, x.data(), n, n, fc);
|
||||
forward(ck.params, x.data(), n, n, fc, fast);
|
||||
backward(ck.params, fc, x.data(), y.data(), n);
|
||||
}
|
||||
t1 = std::chrono::steady_clock::now();
|
||||
|
||||
@ -3,6 +3,6 @@
|
||||
#include "checkpoint.hpp"
|
||||
|
||||
namespace xt {
|
||||
int run_gradcheck(bool verbose);
|
||||
int run_bench(Checkpoint& ck, int iters, int block);
|
||||
int run_gradcheck(bool verbose, bool fast);
|
||||
int run_bench(Checkpoint& ck, int iters, int block, bool fast);
|
||||
}
|
||||
|
||||
@ -204,7 +204,7 @@ void rope_tables(int block, int head_dim, int base, std::vector<float>& cs, std:
|
||||
}
|
||||
}
|
||||
|
||||
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, int seq_len, ForwardCache& fc, bool fast) {
|
||||
const Config& c = p.cfg;
|
||||
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
|
||||
int B = n_tok;
|
||||
@ -226,7 +226,8 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
||||
const float* CS = fc.rope_cos.data();
|
||||
const float* SN = fc.rope_sin.data();
|
||||
|
||||
fc.x.resize(B, d);
|
||||
if (fast) fc.x.resize_fast(B, d);
|
||||
else fc.x.resize(B, d);
|
||||
for (int t = 0; t < B; ++t) {
|
||||
float* xr = fc.x.row(t);
|
||||
const float* er = p.wte.row(tokens[t]);
|
||||
@ -238,7 +239,10 @@ 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];
|
||||
lc.x.resize(B, d);
|
||||
|
||||
if (fast) fc.x.resize_fast(B, d);
|
||||
else fc.x.resize(B, d);
|
||||
|
||||
for (size_t i = 0; i < fc.x.n(); ++i) lc.x.at(i) = fc.x.at(i);
|
||||
|
||||
rmsnorm_forward(lc.x, p.rms_attn[L], c.rms_eps, lc.xb, lc.inv_rms);
|
||||
@ -268,7 +272,8 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
||||
}
|
||||
|
||||
lc.att.resize3d(n_seq * H, S, S);
|
||||
lc.attout.resize(B, d);
|
||||
if (fast) fc.x.resize_fast(B, d);
|
||||
else fc.x.resize(B, d);
|
||||
lc.attout.zero();
|
||||
|
||||
for (int s = 0; s < n_seq; ++s) {
|
||||
@ -307,13 +312,15 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
||||
gemm_nn(lc.attout, p.wo[L], lc.proj);
|
||||
for (size_t i = 0; i < fc.x.n(); ++i) fc.x.at(i) += lc.proj.at(i);
|
||||
|
||||
lc.x2.resize(B, d);
|
||||
if (fast) fc.x.resize_fast(B, d);
|
||||
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);
|
||||
gemm_nn(lc.x2b, p.w1[L], lc.h1);
|
||||
gemm_nn(lc.x2b, p.w3[L], lc.h3);
|
||||
lc.ha.resize(B, f);
|
||||
if (fast) fc.x.resize_fast(B, f);
|
||||
else fc.x.resize(B, f);
|
||||
for (size_t i = 0; i < lc.ha.n(); ++i)
|
||||
lc.ha.at(i) = silu(lc.h1.at(i)) * lc.h3.at(i);
|
||||
gemm_nn(lc.ha, p.w2[L], lc.fo);
|
||||
@ -321,18 +328,19 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
||||
}
|
||||
|
||||
rmsnorm_forward(fc.x, p.rms_final, c.rms_eps, fc.xf, fc.inv_rms);
|
||||
fc.logits.resize(B, c.vocab_out());
|
||||
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);
|
||||
else gemm_nn(fc.xf, p.lm_head, fc.logits);
|
||||
}
|
||||
|
||||
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc) {
|
||||
forward(p, tokens, n_tok, n_tok, fc);
|
||||
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc, bool fast) {
|
||||
forward(p, tokens, n_tok, n_tok, fc, fast);
|
||||
}
|
||||
|
||||
void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits) {
|
||||
void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits, bool fast) {
|
||||
ForwardCache fc;
|
||||
forward(p, tokens, n_tok, n_tok, fc);
|
||||
forward(p, tokens, n_tok, n_tok, fc, fast);
|
||||
const int V = p.cfg.vocab_out();
|
||||
logits.resize(1, V);
|
||||
for (int j = 0; j < V; ++j) logits.at(j) = fc.logits.at((n_tok - 1) * V + j);
|
||||
|
||||
@ -81,10 +81,10 @@ struct ForwardCache {
|
||||
|
||||
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, int seq_len, ForwardCache& fc, bool fast);
|
||||
|
||||
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc);
|
||||
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc, bool fast);
|
||||
|
||||
void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits);
|
||||
void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits, bool fast);
|
||||
|
||||
}
|
||||
@ -71,6 +71,7 @@ struct Tensor {
|
||||
Tensor(int r, int c) : rows(r), cols(c), d((size_t)r * c, 0.0f) {}
|
||||
|
||||
void resize(int r, int c) { rows = r; cols = c; d.assign((size_t)r * c, 0.0f); }
|
||||
void resize_fast(int r, int c) { if (rows != r || cols != c) { rows = r; cols = c; d.assign((size_t)r * c, 0.0f);} else {std::fill(d.begin(), d.end(), 0.0f); }}
|
||||
void resize3d(int h, int r, int c) { rows = r; cols = c; d.assign((size_t)h * r * c, 0.0f); }
|
||||
void zero() { std::fill(d.begin(), d.end(), 0.0f); }
|
||||
int R() const { return rows; }
|
||||
@ -163,6 +164,8 @@ inline float dsilu(float x) {
|
||||
return s * (1.0f + x * (1.0f - s));
|
||||
}
|
||||
|
||||
|
||||
|
||||
inline void rmsnorm_forward(const Tensor& x, const Tensor& w, float eps,
|
||||
Tensor& y, Tensor& inv_rms) {
|
||||
const int N = x.R(), D = x.C();
|
||||
|
||||
@ -29,6 +29,7 @@ struct TrainConfig {
|
||||
int ckpt_every = 0;
|
||||
int val_every = 0;
|
||||
int val_tokens = 20000;
|
||||
bool fast = false;
|
||||
};
|
||||
|
||||
struct Dataset {
|
||||
|
||||
@ -79,7 +79,7 @@ TrainStats train_model(Params& p, const Tokenizer& tok, const std::string& corpu
|
||||
}
|
||||
}
|
||||
|
||||
forward(p, xb.data(), B * block, block, fc);
|
||||
forward(p, xb.data(), B * block, block, fc, tc.fast);
|
||||
float loss = backward(p, fc, xb.data(), yb.data(), B * block);
|
||||
p.adam_step(lr_at(tc, step), tc.beta1, tc.beta2, tc.eps,
|
||||
tc.weight_decay, step + 1, tc.clip);
|
||||
@ -102,7 +102,7 @@ TrainStats train_model(Params& p, const Tokenizer& tok, const std::string& corpu
|
||||
size_t off = (val_ids.size() - block - 1) * w / nwin;
|
||||
std::vector<int> vx(block);
|
||||
for (int t = 0; t < block; ++t) vx[t] = val_ids[off + t];
|
||||
forward(p, vx.data(), block, block, fc);
|
||||
forward(p, vx.data(), block, block, fc, tc.fast);
|
||||
Tensor lg;
|
||||
softmax_eval(p, fc, vx.data(), val_ids.data() + off + 1, block, lg);
|
||||
for (int t = 0; t < block; ++t) vl += lg.at(t);
|
||||
|
||||
@ -1,8 +1,11 @@
|
||||
#pragma once
|
||||
|
||||
#include <sys/ioctl.h>
|
||||
#include <cstdint>
|
||||
#include <memory>
|
||||
#include <optional>
|
||||
#include <termios.h>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
#include <string>
|
||||
|
||||
@ -36,61 +39,88 @@ namespace ui
|
||||
Tui();
|
||||
~Tui();
|
||||
static void update();
|
||||
bool is_running() const;
|
||||
[[nodiscard]] bool is_running() const;
|
||||
private:
|
||||
inline static Mouse mouse = Mouse();
|
||||
inline static std::string chat_input = "";
|
||||
inline static auto need_update_chat = false;
|
||||
inline static auto mouse = Mouse();
|
||||
inline static std::string chat_input;
|
||||
inline static bool cursor = false;
|
||||
inline static uint32_t frame_count = 0;
|
||||
int background = 49;
|
||||
int foreground = 39;
|
||||
inline static int background = 49;
|
||||
inline static int foreground = 39;
|
||||
inline static bool running = true;
|
||||
inline static int width = 80;
|
||||
inline static int height = 24;
|
||||
struct msg_block
|
||||
{
|
||||
std::optional<std::string> left = std::nullopt;
|
||||
std::optional<std::string> right = std::nullopt;
|
||||
};
|
||||
inline static std::vector<msg_block> chat_history;
|
||||
struct SplitDesc {
|
||||
enum Type { VERTICAL, HORIZONTAL };
|
||||
Type type;
|
||||
int position;
|
||||
std::vector<std::unique_ptr<SplitDesc>> children;
|
||||
SplitDesc(Type t, int pos) : type(t), position(pos) {}
|
||||
SplitDesc(const Type t, const int pos) : type(t), position(pos) {}
|
||||
};
|
||||
struct border_charset {
|
||||
std::string_view top_left, top_right, bottom_left, bottom_right;
|
||||
std::string_view bottom, top, horizontal, vertical, bg;
|
||||
std::string_view cross, tee_up, tee_down, tee_left, tee_right;
|
||||
};
|
||||
struct scroll_chars
|
||||
{
|
||||
std::string_view top;
|
||||
std::string_view full;
|
||||
std::string_view bottom;
|
||||
};
|
||||
struct TuiCell {
|
||||
std::string ch;
|
||||
int width;
|
||||
TuiCell() : ch(" "), width(1) {}
|
||||
TuiCell(const std::string& c, int w) : ch(c), width(w) {}
|
||||
TuiCell(std::string c, const int w) : ch(std::move(c)), width(w) {}
|
||||
};
|
||||
std::string proces_chars = "⠙⠸⠴⠦⠇⠋";
|
||||
inline static struct termios orig_termios;
|
||||
inline static const border_charset border_clip {
|
||||
inline static constexpr scroll_chars main_scroll_charset {
|
||||
.top = "▗",
|
||||
.full = "▐",
|
||||
.bottom = "▝",
|
||||
};
|
||||
inline static constexpr border_charset border_clip {
|
||||
.top_left="╭", .top_right="╮", .bottom_left="╰", .bottom_right="╯",
|
||||
.bottom="─", .top="─", .horizontal="─", .vertical="│", .bg=" ",
|
||||
.cross="┼", .tee_up="┴", .tee_down="┬", .tee_left="┤", .tee_right="├"
|
||||
};
|
||||
inline static const border_charset border_ascii {
|
||||
inline static constexpr border_charset border_ascii {
|
||||
.top_left="+", .top_right="+", .bottom_left="+", .bottom_right="+",
|
||||
.bottom="-", .top="-", .vertical="|", .bg=" ", .cross="+",
|
||||
.tee_up="+", .tee_down="+", .tee_left="+", .tee_right="+"
|
||||
.bottom="-", .top="-", .horizontal="-", .vertical="|", .bg=" ",
|
||||
.cross="+", .tee_up="+", .tee_down="+", .tee_left="+",
|
||||
.tee_right="+"
|
||||
};
|
||||
inline static const border_charset border_bold {
|
||||
inline static constexpr border_charset border_bold {
|
||||
.top_left="┏", .top_right="┓", .bottom_left="┗", .bottom_right="┛",
|
||||
.bottom="━", .top="━", .vertical="┃", .bg=" "
|
||||
.bottom="━", .top="━", .horizontal="━", .vertical="┃",
|
||||
.bg=" ", .cross = "╋", .tee_up = "┻", .tee_down = "┳", .tee_left = "┨",
|
||||
.tee_right = "┣"
|
||||
};
|
||||
inline static const border_charset border_solid {
|
||||
inline static constexpr border_charset border_solid {
|
||||
.top_left="▟", .top_right="▙", .bottom_left="▜", .bottom_right="▛",
|
||||
.bottom="█", .top="█", .vertical="█", .bg=" "
|
||||
.bottom="█", .top="█", .horizontal = "█", .vertical="█",
|
||||
.bg=" ", .cross = "█", .tee_up = "█", .tee_down = "█", .tee_left = "█",
|
||||
.tee_right = "█",
|
||||
};
|
||||
inline static const border_charset border_btn {
|
||||
inline static constexpr border_charset border_btn {
|
||||
.top_left="▟", .top_right="▙", .bottom_left="▜", .bottom_right="▛",
|
||||
.bottom="█", .top="█", .vertical="█", .bg="█"
|
||||
.bottom="█", .top="█", .horizontal = "█", .vertical="█",
|
||||
.bg="█", .cross = "█", .tee_up = "█", .tee_down = "█", .tee_left = "█",
|
||||
.tee_right = "█"
|
||||
};
|
||||
static void get_console_size();
|
||||
static void init_draw();
|
||||
static winsize get_console_size();
|
||||
static bool wait_for_input_or_timeout(int timeout_ms);
|
||||
void apply_color() const;
|
||||
static void apply_color();
|
||||
static void disable_raw_mode();
|
||||
static void enable_raw_mode();
|
||||
static void move_cursor(int x, int y);
|
||||
@ -98,19 +128,20 @@ namespace ui
|
||||
static void draw_border(int w, int h, border_charset charset = border_clip, int x = 1, int y = 1);
|
||||
static void hide_cursor();
|
||||
static void show_cursor();
|
||||
void set_background(Color clr);
|
||||
void set_foreground(Color clr);
|
||||
static void set_background(Color clr);
|
||||
static void set_foreground(Color clr);
|
||||
static void move_cursor_right(int n);
|
||||
static void move_cursor_x(int col);
|
||||
static void move_cursor_y(int row);
|
||||
static void invert_color();
|
||||
static void draw_button(int x, int y, int w, int h, std::string_view label, border_charset charset = border_solid);
|
||||
static void invert_color_off();
|
||||
void draw_split_border(int w, int h, border_charset charset, int x, int y, const SplitDesc* root);
|
||||
static void draw_split_border(int w, int h, border_charset charset, int x, int y, const SplitDesc* root);
|
||||
static void apply_splits_recursive(std::vector<std::vector<TuiCell>>& grid, int offset_x, int offset_y, int w, int h, border_charset charset, const SplitDesc* split);
|
||||
static void apply_single_split(std::vector<std::vector<TuiCell>>& grid, int offset_x, int offset_y, int w, int h, border_charset charset, const SplitDesc* split);
|
||||
static void draw_border_to_buffer(std::vector<std::vector<TuiCell>>& grid, int w, int h, border_charset charset);
|
||||
static void push_msg_left(const std::string& msg);
|
||||
static void push_msg_right(const std::string& msg);
|
||||
static void draw_msg();
|
||||
};
|
||||
|
||||
Tui app = Tui();
|
||||
}
|
||||
@ -1,7 +1,68 @@
|
||||
#include "message.h"
|
||||
#include <iostream>
|
||||
#include <ostream>
|
||||
|
||||
#include "./../tui.hpp"
|
||||
|
||||
namespace ui
|
||||
{
|
||||
void Tui::push_msg_left(const std::string& msg)
|
||||
{
|
||||
if (!chat_history.empty() && !chat_history.back().left.has_value())
|
||||
{
|
||||
chat_history.back().left = msg;
|
||||
}
|
||||
else
|
||||
{
|
||||
msg_block block = {
|
||||
.left = msg,
|
||||
.right = std::nullopt
|
||||
};
|
||||
chat_history.push_back(block);
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
void Tui::push_msg_right(const std::string& msg)
|
||||
{
|
||||
if (!chat_history.empty() && !chat_history.back().right.has_value())
|
||||
{
|
||||
chat_history.back().right = msg;
|
||||
}
|
||||
else
|
||||
{
|
||||
msg_block block = {
|
||||
.left = std::nullopt,
|
||||
.right = msg,
|
||||
};
|
||||
chat_history.push_back(block);
|
||||
}
|
||||
}
|
||||
|
||||
void Tui::draw_msg()
|
||||
{
|
||||
const auto root = std::make_unique<SplitDesc>(SplitDesc::HORIZONTAL, 2);
|
||||
|
||||
int total_hg = chat_history.size() * 5;
|
||||
|
||||
int offsetY = 0;
|
||||
for (int q = 0; q < chat_history.size(); q++)
|
||||
{
|
||||
int wh = std::max(3, static_cast<int>(chat_history[q].right.value_or("").length()));
|
||||
draw_split_border(wh + 2, 5, border_clip, width - wh - 2, 2 + offsetY, root.get());
|
||||
move_cursor(width - wh - 1, 3 + offsetY);
|
||||
std::cout << "You" << std::flush;
|
||||
move_cursor(width - wh - 1, 5 + offsetY);
|
||||
std::cout << chat_history[q].right.value_or("") << std::flush;
|
||||
offsetY += 5;
|
||||
|
||||
|
||||
wh = std::max(22, static_cast<int>(chat_history[q].left.value_or("").length()));
|
||||
draw_split_border(wh + 2, 5, border_clip, (width / 100) * 27 + 2, 2 + offsetY, root.get());
|
||||
move_cursor((width / 100) * 27 + 3, 3 + offsetY);
|
||||
std::cout << "Ai (Processing prompt█)" << std::flush;
|
||||
move_cursor((width / 100) * 27 + 3, 5 + offsetY);
|
||||
std::cout << chat_history[q].left.value_or("") << std::flush;
|
||||
offsetY += 5;
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
@ -3,7 +3,6 @@
|
||||
#include <chrono>
|
||||
#include <csignal>
|
||||
#include <cstring>
|
||||
#include <sys/ioctl.h>
|
||||
#include <unistd.h>
|
||||
#include <sys/select.h>
|
||||
#include <iostream>
|
||||
@ -13,16 +12,23 @@ namespace ui
|
||||
{
|
||||
Tui::Tui()
|
||||
{
|
||||
get_console_size();
|
||||
int panel_w = (width / 100) * 27;
|
||||
std::cout << "Ширина (колонок): " << width << "\n";
|
||||
std::cout << "Высота (строк): " << height << "\n";
|
||||
init_draw();
|
||||
|
||||
enable_raw_mode();
|
||||
}
|
||||
|
||||
void Tui::init_draw()
|
||||
{
|
||||
winsize wins = get_console_size();
|
||||
width = wins.ws_col;
|
||||
height = wins.ws_row;
|
||||
const int panel_w = (width / 100) * 27;
|
||||
clear();
|
||||
hide_cursor();
|
||||
set_foreground(Color::White);
|
||||
apply_color();
|
||||
|
||||
auto root = std::make_unique<SplitDesc>(SplitDesc::VERTICAL, (width / 100) * 27);
|
||||
const auto root = std::make_unique<SplitDesc>(SplitDesc::VERTICAL, (width / 100) * 27);
|
||||
|
||||
root->children.push_back(nullptr);
|
||||
|
||||
@ -41,8 +47,6 @@ namespace ui
|
||||
|
||||
move_cursor(width - 8, height - 1);
|
||||
std::cout << "[enter>" <<std::flush;
|
||||
|
||||
enable_raw_mode();
|
||||
}
|
||||
|
||||
Tui::~Tui()
|
||||
@ -122,9 +126,23 @@ namespace ui
|
||||
}
|
||||
i++;
|
||||
}
|
||||
else if (c == 13 || c == 10) {
|
||||
if (!chat_input.empty()) {
|
||||
push_msg_right(chat_input);
|
||||
push_msg_left("Hi! How can I help you today?");
|
||||
chat_input.clear();
|
||||
input_changed = true;
|
||||
draw_msg();
|
||||
}
|
||||
i++;
|
||||
}
|
||||
else if (c >= 32) {
|
||||
chat_input += c;
|
||||
input_changed = true;
|
||||
int spaces_to_draw = std::max(0, (width - ((width / 100) * 27) - 15 - static_cast<int>(chat_input.length())));
|
||||
if (spaces_to_draw > 0)
|
||||
{
|
||||
chat_input += c;
|
||||
input_changed = true;
|
||||
}
|
||||
i++;
|
||||
}
|
||||
else {
|
||||
@ -143,11 +161,22 @@ namespace ui
|
||||
} else {
|
||||
cursor = !cursor;
|
||||
}
|
||||
|
||||
winsize wins = get_console_size();
|
||||
if (width != wins.ws_col || height != wins.ws_row)
|
||||
{
|
||||
init_draw();
|
||||
width = wins.ws_col;
|
||||
height = wins.ws_row;
|
||||
}
|
||||
|
||||
last_blink = now;
|
||||
|
||||
int input_x = ((width / 100) * 27) + 3;
|
||||
move_cursor(input_x, height - 1);
|
||||
std::cout << chat_input << (cursor ? "_ " : " ") << std::flush;
|
||||
int spaces_to_draw = std::max(0, (width - ((width / 100) * 27) - 15 - static_cast<int>(chat_input.length())));
|
||||
std::cout << chat_input << (cursor ? "_ " : " ") << std::string(spaces_to_draw, ' ') << std::flush;
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
@ -156,17 +185,14 @@ namespace ui
|
||||
return running;
|
||||
}
|
||||
|
||||
void Tui::get_console_size()
|
||||
winsize Tui::get_console_size()
|
||||
{
|
||||
winsize w;
|
||||
|
||||
if (ioctl(STDOUT_FILENO, TIOCGWINSZ, &w) == 0) {
|
||||
width = w.ws_col;
|
||||
height = w.ws_row;
|
||||
}
|
||||
ioctl(STDOUT_FILENO, TIOCGWINSZ, &w);
|
||||
return w;
|
||||
}
|
||||
|
||||
void Tui::apply_color() const
|
||||
void Tui::apply_color()
|
||||
{
|
||||
std::cout << "\033["<< foreground << ";" << background <<"m" << std::flush;
|
||||
}
|
||||
@ -264,7 +290,6 @@ namespace ui
|
||||
|
||||
void Tui::draw_split_border(int w, int h, border_charset charset, int x, int y, const SplitDesc* root)
|
||||
{
|
||||
// Создаем сетку из TuiCell, инициализированную фоновым символом
|
||||
std::vector<std::vector<TuiCell>> grid(h, std::vector<TuiCell>(w, TuiCell(std::string(charset.bg), 1)));
|
||||
|
||||
draw_border_to_buffer(grid, w, h, charset);
|
||||
@ -281,7 +306,7 @@ namespace ui
|
||||
std::cout << grid[i][j].ch;
|
||||
}
|
||||
if (i < h - 1) {
|
||||
std::cout << "\n";
|
||||
std::cout << "\033[1B\033[" << x << "G" << std::flush;
|
||||
}
|
||||
}
|
||||
|
||||
@ -321,42 +346,50 @@ namespace ui
|
||||
|
||||
void Tui::apply_single_split(std::vector<std::vector<TuiCell>>& grid, int offset_x, int offset_y, int w, int h, border_charset charset, const SplitDesc* split)
|
||||
{
|
||||
if (grid.empty()) {
|
||||
return;
|
||||
}
|
||||
if (split->type == SplitDesc::VERTICAL) {
|
||||
int x = offset_x + split->position;
|
||||
|
||||
for (int i = 0; i < h; i++) {
|
||||
int y = offset_y + i;
|
||||
TuiCell& cell = grid[y][x];
|
||||
|
||||
if (i == 0) {
|
||||
if (cell.ch == charset.top) cell.ch = std::string(charset.tee_down);
|
||||
} else if (i == h - 1) {
|
||||
if (cell.ch == charset.bottom) cell.ch = std::string(charset.tee_up);
|
||||
} else {
|
||||
if (cell.ch == charset.bg) {
|
||||
cell.ch = std::string(charset.vertical);
|
||||
} else if (cell.ch == charset.horizontal || cell.ch == charset.top || cell.ch == charset.bottom) {
|
||||
cell.ch = std::string(charset.cross);
|
||||
if (!grid[0].empty() && x >= 0 && x < static_cast<int>(grid[0].size())) {
|
||||
for (int i = 0; i < h; i++) {
|
||||
int y = offset_y + i;
|
||||
if (y >= 0 && y < static_cast<int>(grid.size())) {
|
||||
if (x < static_cast<int>(grid[y].size())) {
|
||||
TuiCell& cell = grid[y][x];
|
||||
if (i == 0) {
|
||||
if (cell.ch == charset.top) cell.ch = std::string(charset.tee_down);
|
||||
} else if (i == h - 1) {
|
||||
if (cell.ch == charset.bottom) cell.ch = std::string(charset.tee_up);
|
||||
} else {
|
||||
if (cell.ch == charset.bg) {
|
||||
cell.ch = std::string(charset.vertical);
|
||||
} else if (cell.ch == charset.horizontal || cell.ch == charset.top || cell.ch == charset.bottom) {
|
||||
cell.ch = std::string(charset.cross);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} else { // HORIZONTAL
|
||||
} else {
|
||||
int y = offset_y + split->position;
|
||||
|
||||
for (int j = 0; j < w; j++) {
|
||||
int x = offset_x + j;
|
||||
TuiCell& cell = grid[y][x];
|
||||
|
||||
if (j == 0) {
|
||||
if (cell.ch == charset.vertical) cell.ch = std::string(charset.tee_right);
|
||||
} else if (j == w - 1) {
|
||||
if (cell.ch == charset.vertical) cell.ch = std::string(charset.tee_left);
|
||||
} else {
|
||||
if (cell.ch == charset.bg) {
|
||||
// Используем horizontal вместо bottom для семантической корректности горизонтального разделителя
|
||||
cell.ch = std::string(charset.horizontal);
|
||||
} else if (cell.ch == charset.vertical) {
|
||||
cell.ch = std::string(charset.cross);
|
||||
if (y >= 0 && y < static_cast<int>(grid.size())) {
|
||||
for (int j = 0; j < w; j++) {
|
||||
int x = offset_x + j;
|
||||
if (x >= 0 && x < static_cast<int>(grid[y].size())) {
|
||||
TuiCell& cell = grid[y][x];
|
||||
if (j == 0) {
|
||||
if (cell.ch == charset.vertical) cell.ch = std::string(charset.tee_right);
|
||||
} else if (j == w - 1) {
|
||||
if (cell.ch == charset.vertical) cell.ch = std::string(charset.tee_left);
|
||||
} else {
|
||||
if (cell.ch == charset.bg) {
|
||||
cell.ch = std::string(charset.horizontal);
|
||||
} else if (cell.ch == charset.vertical) {
|
||||
cell.ch = std::string(charset.cross);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user