--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 конфиг с настройками и списком моделей
|
--conf NAME.conf конфиг с настройками и списком моделей
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### gradcheck - проверить градиенты численно
|
||||||
|
|
||||||
|
```
|
||||||
|
--fast эксперементаль флаг который ускоряет foreword примерно в x2-x5
|
||||||
|
```
|
||||||
|
|
||||||
|
### bench - замерить скорость форварда и шага обучения
|
||||||
|
|
||||||
|
```
|
||||||
|
--fast эксперементаль флаг который ускоряет foreword примерно в x2-x5
|
||||||
|
```
|
||||||
|
|
||||||
### update - проверка обновления
|
### update - проверка обновления
|
||||||
|
|
||||||
### info / gradcheck / bench
|
### info
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
Xenith info models/my.xnh # конфигурация и словарь
|
Xenith info models/my.xnh # конфигурация и словарь
|
||||||
Xenith gradcheck # градиенты против численных
|
|
||||||
Xenith bench models/my.xnh # скорость
|
|
||||||
```
|
```
|
||||||
|
|
||||||
`gradcheck` стоит запускать после любой правки в `model.cpp` / `backward.cpp` —
|
`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-every N валидация каждые N шагов (0 = нет)\n" \
|
||||||
" --val-tokens N размер валидации (по умолчанию 20000)\n" \
|
" --val-tokens N размер валидации (по умолчанию 20000)\n" \
|
||||||
" --threads N потоков (0 = все ядра)\n" \
|
" --threads N потоков (0 = все ядра)\n" \
|
||||||
" --resume продолжить с сохранённого step (Adam-state из файла)\n"
|
" --resume продолжить с сохранённого step (Adam-state из файла)\n" \
|
||||||
|
" --fast эксперементаль флаг который ускоряет foreword примерно в x2-x5\n"
|
||||||
#else
|
#else
|
||||||
#define OPT_TRAIN ""
|
#define OPT_TRAIN ""
|
||||||
#endif
|
#endif
|
||||||
@ -109,6 +110,18 @@
|
|||||||
#else
|
#else
|
||||||
#define OPT_GEN ""
|
#define OPT_GEN ""
|
||||||
#endif
|
#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
|
#ifdef OLLAMA
|
||||||
#define OPT_OLLAMA "\nОПЦИИ ollama\n" \
|
#define OPT_OLLAMA "\nОПЦИИ ollama\n" \
|
||||||
" --port PORT порт на котором будет открыт сервер\n" \
|
" --port PORT порт на котором будет открыт сервер\n" \
|
||||||
@ -133,6 +146,9 @@
|
|||||||
OPT_NEW \
|
OPT_NEW \
|
||||||
OPT_TRAIN \
|
OPT_TRAIN \
|
||||||
OPT_GEN \
|
OPT_GEN \
|
||||||
OPT_OLLAMA
|
OPT_GRADCHECK \
|
||||||
|
OPT_BENCH \
|
||||||
|
OPT_OLLAMA \
|
||||||
|
"\n"
|
||||||
|
|
||||||
#endif
|
#endif
|
||||||
51
main.cpp
51
main.cpp
@ -1,7 +1,7 @@
|
|||||||
#include "xenith/core/checkpoint.hpp"
|
#include "xenith/core/checkpoint.hpp"
|
||||||
#include "xenith/core/generate.hpp"
|
#include "xenith/core/generate.hpp"
|
||||||
#include "xenith/core/train.hpp"
|
#include "xenith/core/train.hpp"
|
||||||
#include "xenith/core/gradcheck.h"
|
#include "xenith/core/gradcheck.hpp"
|
||||||
#include "xenith/ui/tui.hpp"
|
#include "xenith/ui/tui.hpp"
|
||||||
#include <cstdio>
|
#include <cstdio>
|
||||||
#include "help.h"
|
#include "help.h"
|
||||||
@ -149,6 +149,7 @@ static int cmd_train(const Args& a) {
|
|||||||
tc.val_every = a.geti("val-every", 0);
|
tc.val_every = a.geti("val-every", 0);
|
||||||
tc.val_tokens = a.geti("val-tokens", 20000);
|
tc.val_tokens = a.geti("val-tokens", 20000);
|
||||||
tc.threads = a.geti("threads", 0);
|
tc.threads = a.geti("threads", 0);
|
||||||
|
tc.fast = a.has("fast");
|
||||||
|
|
||||||
if (tc.threads > 0) set_threads(tc.threads);
|
if (tc.threads > 0) set_threads(tc.threads);
|
||||||
if (tc.block > ck.params.cfg.block_size) tc.block = ck.params.cfg.block_size;
|
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) {
|
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) {
|
static int cmd_bench(const Args& a) {
|
||||||
@ -237,7 +238,7 @@ static int cmd_bench(const Args& a) {
|
|||||||
Checkpoint ck;
|
Checkpoint ck;
|
||||||
std::string err;
|
std::string err;
|
||||||
if (!load_checkpoint(a.pos[0], ck, err)) { fprintf(stderr, "bench: %s\n", err.c_str()); return 1; }
|
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) {
|
static int cmd_ollama(const Args& a) {
|
||||||
@ -302,10 +303,13 @@ static bool available_new_version() {
|
|||||||
Version remote_ver = parse_version(remote_str);
|
Version remote_ver = parse_version(remote_str);
|
||||||
Version local_ver = parse_version(local_str);
|
Version local_ver = parse_version(local_str);
|
||||||
|
|
||||||
|
if (is_newer(remote_ver, local_ver))
|
||||||
|
{
|
||||||
std::cout << "Локальная версия: " << local_ver.major << "."
|
std::cout << "Локальная версия: " << local_ver.major << "."
|
||||||
<< local_ver.minor << "." << local_ver.patch << std::endl;
|
<< local_ver.minor << "." << local_ver.patch << std::endl;
|
||||||
std::cout << "Удалённая версия: " << remote_ver.major << "."
|
std::cout << "Удалённая версия: " << remote_ver.major << "."
|
||||||
<< remote_ver.minor << "." << remote_ver.patch << std::endl;
|
<< remote_ver.minor << "." << remote_ver.patch << std::endl;
|
||||||
|
}
|
||||||
|
|
||||||
return is_newer(remote_ver, local_ver);
|
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 REPO_URL = "https://git.bipfr.ru/BIPfR/Xenith.git";
|
||||||
const std::string TMP_CLONE = "/tmp/xenith_update";
|
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;
|
std::cout << "Скачивание новой версии..." << std::endl;
|
||||||
int result = system(("git clone --depth 1 " + REPO_URL + " " + TMP_CLONE).c_str());
|
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());
|
result = system(("cp -ra " + TMP_CLONE + "/. .").c_str());
|
||||||
if (result != 0) {
|
if (result != 0) {
|
||||||
std::cerr << "Ошибка копирования файлов" << std::endl;
|
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;
|
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;
|
std::cout << "Обновление завершено!" << std::endl;
|
||||||
return 0;
|
return 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
static int cmd_tui(const Args& a)
|
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;
|
std::cout << "\033[2J\033[H" << std::flush;
|
||||||
return 0;
|
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;
|
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) {
|
if (map_size != mdl.config.vocab_size) {
|
||||||
std::cerr << "Warning: vocab size mismatch! "
|
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) {
|
if (magic != XNH_MAGIC) {
|
||||||
char buf[64];
|
char buf[64];
|
||||||
std::snprintf(buf, sizeof buf,
|
std::snprintf(buf, sizeof buf,
|
||||||
"плохая магия 0x%08X (ожидалась 0x%08X) — это не Xenith-модель",
|
"Файл не является поддерживаемым",
|
||||||
magic, XNH_MAGIC);
|
magic, XNH_MAGIC);
|
||||||
err = buf;
|
err = buf;
|
||||||
return false;
|
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,
|
void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
||||||
const std::vector<float>& cs, const std::vector<float>& sn,
|
const std::vector<float>& cs, const std::vector<float>& sn,
|
||||||
Tensor& logits) {
|
Tensor& logits, bool fast) {
|
||||||
const Config& c = p.cfg;
|
const Config& c = p.cfg;
|
||||||
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
|
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);
|
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); }
|
for (int j = 0; j < hd; ++j) { kp[j] = k.at(h * hd + j); vp[j] = v.at(h * hd + j); }
|
||||||
}
|
}
|
||||||
|
|
||||||
// attention по всем позициям 0..pos
|
if (fast) attout.resize_fast(1, d);
|
||||||
attout.resize(1, d);
|
else attout.resize(1, d);
|
||||||
attout.zero();
|
attout.zero();
|
||||||
for (int h = 0; h < H; ++h) {
|
for (int h = 0; h < H; ++h) {
|
||||||
const float* qr = q.row(0) + h * hd;
|
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);
|
rmsnorm_forward(x, p.rms_ffn[L], c.rms_eps, x2b, ln);
|
||||||
gemm_nn(x2b, p.w1[L], h1);
|
gemm_nn(x2b, p.w1[L], h1);
|
||||||
gemm_nn(x2b, p.w3[L], h3);
|
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);
|
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);
|
gemm_nn(ha, p.w2[L], fo);
|
||||||
for (int j = 0; j < d; ++j) x.at(j) += fo.at(j);
|
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);
|
||||||
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);
|
if (p.tied()) gemm_nt(lnf, p.wte, logits);
|
||||||
else gemm_nn(lnf, p.lm_head, 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,
|
std::string generate(const Params& p, const Tokenizer& tok,
|
||||||
const std::string& prompt, const GenConfig& gc,
|
const std::string& prompt, const GenConfig& gc,
|
||||||
std::vector<int>* out_ids) {
|
std::vector<int>* out_ids, bool fast) {
|
||||||
const Config& c = p.cfg;
|
const Config& c = p.cfg;
|
||||||
KvCache kv;
|
KvCache kv;
|
||||||
kv.init(c.n_layer, c.block_size, c.n_head, c.head_dim());
|
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);
|
ids.resize(keep);
|
||||||
kv.reset();
|
kv.reset();
|
||||||
for (int t = 0; t < keep; ++t) {
|
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++;
|
kv.len++;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const int last = ids.back();
|
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++;
|
kv.len++;
|
||||||
|
|
||||||
int next = sample(logits, gc, rng, ids);
|
int next = sample(logits, gc, rng, ids);
|
||||||
|
|||||||
@ -18,7 +18,7 @@ struct GenConfig {
|
|||||||
|
|
||||||
std::string generate(const Params& p, const Tokenizer& tok,
|
std::string generate(const Params& p, const Tokenizer& tok,
|
||||||
const std::string& prompt, const GenConfig& gc,
|
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,
|
std::vector<std::string> generate_batch(const Params& p, const Tokenizer& tok,
|
||||||
const std::vector<std::string>& prompts,
|
const std::vector<std::string>& prompts,
|
||||||
|
|||||||
@ -1,6 +1,6 @@
|
|||||||
// gradcheck.cpp — проверяем, что обратный проход совпадает с численным
|
// gradcheck.cpp — проверяем, что обратный проход совпадает с численным
|
||||||
// и заодно меряем скорость.
|
// и заодно меряем скорость.
|
||||||
#include "gradcheck.h"
|
#include "gradcheck.hpp"
|
||||||
#include "generate.hpp"
|
#include "generate.hpp"
|
||||||
#include "train.hpp"
|
#include "train.hpp"
|
||||||
#include <cstdio>
|
#include <cstdio>
|
||||||
@ -14,9 +14,9 @@ namespace xt {
|
|||||||
namespace {
|
namespace {
|
||||||
|
|
||||||
// loss на фиксированном батче, без обратного прохода
|
// 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;
|
ForwardCache fc;
|
||||||
forward(p, x.data(), n, fc);
|
forward(p, x.data(), n, fc, fast);
|
||||||
Tensor loss;
|
Tensor loss;
|
||||||
return softmax_eval(p, fc, x.data(), y.data(), n, 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;
|
Config c;
|
||||||
c.vocab_size = 24;
|
c.vocab_size = 24;
|
||||||
c.n_layer = 2;
|
c.n_layer = 2;
|
||||||
@ -53,7 +53,7 @@ int run_gradcheck(bool verbose) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
ForwardCache fc;
|
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);
|
backward(p, fc, x.data(), y.data(), n);
|
||||||
|
|
||||||
auto tensors = p.all_w();
|
auto tensors = p.all_w();
|
||||||
@ -80,9 +80,9 @@ int run_gradcheck(bool verbose) {
|
|||||||
const float orig = W.at(i);
|
const float orig = W.at(i);
|
||||||
|
|
||||||
W.at(i) = orig + eps;
|
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;
|
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;
|
W.at(i) = orig;
|
||||||
|
|
||||||
const float num = (lp - lm) / (2.0f * eps);
|
const float num = (lp - lm) / (2.0f * eps);
|
||||||
@ -112,7 +112,7 @@ int run_gradcheck(bool verbose) {
|
|||||||
return 1;
|
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;
|
const Config& c = ck.params.cfg;
|
||||||
if (block > c.block_size) block = c.block_size;
|
if (block > c.block_size) block = c.block_size;
|
||||||
const int n = block;
|
const int n = block;
|
||||||
@ -129,18 +129,18 @@ int run_bench(Checkpoint& ck, int iters, int block) {
|
|||||||
|
|
||||||
// прогрев
|
// прогрев
|
||||||
ForwardCache fc;
|
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);
|
backward(ck.params, fc, x.data(), y.data(), n);
|
||||||
|
|
||||||
auto t0 = std::chrono::steady_clock::now();
|
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();
|
auto t1 = std::chrono::steady_clock::now();
|
||||||
double d_fwd = std::chrono::duration<double>(t1 - t0).count();
|
double d_fwd = std::chrono::duration<double>(t1 - t0).count();
|
||||||
float fwd_tps = (double)n * iters / d_fwd;
|
float fwd_tps = (double)n * iters / d_fwd;
|
||||||
|
|
||||||
t0 = std::chrono::steady_clock::now();
|
t0 = std::chrono::steady_clock::now();
|
||||||
for (int i = 0; i < iters; ++i) {
|
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);
|
backward(ck.params, fc, x.data(), y.data(), n);
|
||||||
}
|
}
|
||||||
t1 = std::chrono::steady_clock::now();
|
t1 = std::chrono::steady_clock::now();
|
||||||
|
|||||||
@ -3,6 +3,6 @@
|
|||||||
#include "checkpoint.hpp"
|
#include "checkpoint.hpp"
|
||||||
|
|
||||||
namespace xt {
|
namespace xt {
|
||||||
int run_gradcheck(bool verbose);
|
int run_gradcheck(bool verbose, bool fast);
|
||||||
int run_bench(Checkpoint& ck, int iters, int block);
|
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 Config& c = p.cfg;
|
||||||
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
|
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
|
||||||
int B = n_tok;
|
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* CS = fc.rope_cos.data();
|
||||||
const float* SN = fc.rope_sin.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) {
|
for (int t = 0; t < B; ++t) {
|
||||||
float* xr = fc.x.row(t);
|
float* xr = fc.x.row(t);
|
||||||
const float* er = p.wte.row(tokens[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) {
|
for (int L = 0; L < c.n_layer; ++L) {
|
||||||
LayerCache& lc = fc.layers[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);
|
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);
|
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.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();
|
lc.attout.zero();
|
||||||
|
|
||||||
for (int s = 0; s < n_seq; ++s) {
|
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);
|
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);
|
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);
|
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);
|
||||||
gemm_nn(lc.x2b, p.w1[L], lc.h1);
|
gemm_nn(lc.x2b, p.w1[L], lc.h1);
|
||||||
gemm_nn(lc.x2b, p.w3[L], lc.h3);
|
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)
|
for (size_t i = 0; i < lc.ha.n(); ++i)
|
||||||
lc.ha.at(i) = silu(lc.h1.at(i)) * lc.h3.at(i);
|
lc.ha.at(i) = silu(lc.h1.at(i)) * lc.h3.at(i);
|
||||||
gemm_nn(lc.ha, p.w2[L], lc.fo);
|
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);
|
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);
|
if (p.tied()) gemm_nt(fc.xf, p.wte, fc.logits);
|
||||||
else gemm_nn(fc.xf, p.lm_head, 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) {
|
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc, bool fast) {
|
||||||
forward(p, tokens, n_tok, n_tok, fc);
|
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;
|
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();
|
const int V = p.cfg.vocab_out();
|
||||||
logits.resize(1, V);
|
logits.resize(1, V);
|
||||||
for (int j = 0; j < V; ++j) logits.at(j) = fc.logits.at((n_tok - 1) * V + j);
|
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,
|
void rope_tables(int block, int head_dim, int base,
|
||||||
std::vector<float>& cs, std::vector<float>& sn);
|
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) {}
|
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(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 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); }
|
void zero() { std::fill(d.begin(), d.end(), 0.0f); }
|
||||||
int R() const { return rows; }
|
int R() const { return rows; }
|
||||||
@ -163,6 +164,8 @@ inline float dsilu(float x) {
|
|||||||
return s * (1.0f + x * (1.0f - s));
|
return s * (1.0f + x * (1.0f - s));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
inline void rmsnorm_forward(const Tensor& x, const Tensor& w, float eps,
|
inline void rmsnorm_forward(const Tensor& x, const Tensor& w, float eps,
|
||||||
Tensor& y, Tensor& inv_rms) {
|
Tensor& y, Tensor& inv_rms) {
|
||||||
const int N = x.R(), D = x.C();
|
const int N = x.R(), D = x.C();
|
||||||
|
|||||||
@ -29,6 +29,7 @@ struct TrainConfig {
|
|||||||
int ckpt_every = 0;
|
int ckpt_every = 0;
|
||||||
int val_every = 0;
|
int val_every = 0;
|
||||||
int val_tokens = 20000;
|
int val_tokens = 20000;
|
||||||
|
bool fast = false;
|
||||||
};
|
};
|
||||||
|
|
||||||
struct Dataset {
|
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);
|
float loss = backward(p, fc, xb.data(), yb.data(), B * block);
|
||||||
p.adam_step(lr_at(tc, step), tc.beta1, tc.beta2, tc.eps,
|
p.adam_step(lr_at(tc, step), tc.beta1, tc.beta2, tc.eps,
|
||||||
tc.weight_decay, step + 1, tc.clip);
|
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;
|
size_t off = (val_ids.size() - block - 1) * w / nwin;
|
||||||
std::vector<int> vx(block);
|
std::vector<int> vx(block);
|
||||||
for (int t = 0; t < block; ++t) vx[t] = val_ids[off + t];
|
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;
|
Tensor lg;
|
||||||
softmax_eval(p, fc, vx.data(), val_ids.data() + off + 1, block, 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);
|
for (int t = 0; t < block; ++t) vl += lg.at(t);
|
||||||
|
|||||||
@ -1,8 +1,11 @@
|
|||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
|
#include <sys/ioctl.h>
|
||||||
#include <cstdint>
|
#include <cstdint>
|
||||||
#include <memory>
|
#include <memory>
|
||||||
|
#include <optional>
|
||||||
#include <termios.h>
|
#include <termios.h>
|
||||||
|
#include <utility>
|
||||||
#include <vector>
|
#include <vector>
|
||||||
#include <string>
|
#include <string>
|
||||||
|
|
||||||
@ -36,61 +39,88 @@ namespace ui
|
|||||||
Tui();
|
Tui();
|
||||||
~Tui();
|
~Tui();
|
||||||
static void update();
|
static void update();
|
||||||
bool is_running() const;
|
[[nodiscard]] bool is_running() const;
|
||||||
private:
|
private:
|
||||||
inline static Mouse mouse = Mouse();
|
inline static auto need_update_chat = false;
|
||||||
inline static std::string chat_input = "";
|
inline static auto mouse = Mouse();
|
||||||
|
inline static std::string chat_input;
|
||||||
inline static bool cursor = false;
|
inline static bool cursor = false;
|
||||||
inline static uint32_t frame_count = 0;
|
inline static uint32_t frame_count = 0;
|
||||||
int background = 49;
|
inline static int background = 49;
|
||||||
int foreground = 39;
|
inline static int foreground = 39;
|
||||||
inline static bool running = true;
|
inline static bool running = true;
|
||||||
inline static int width = 80;
|
inline static int width = 80;
|
||||||
inline static int height = 24;
|
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 {
|
struct SplitDesc {
|
||||||
enum Type { VERTICAL, HORIZONTAL };
|
enum Type { VERTICAL, HORIZONTAL };
|
||||||
Type type;
|
Type type;
|
||||||
int position;
|
int position;
|
||||||
std::vector<std::unique_ptr<SplitDesc>> children;
|
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 {
|
struct border_charset {
|
||||||
std::string_view top_left, top_right, bottom_left, bottom_right;
|
std::string_view top_left, top_right, bottom_left, bottom_right;
|
||||||
std::string_view bottom, top, horizontal, vertical, bg;
|
std::string_view bottom, top, horizontal, vertical, bg;
|
||||||
std::string_view cross, tee_up, tee_down, tee_left, tee_right;
|
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 {
|
struct TuiCell {
|
||||||
std::string ch;
|
std::string ch;
|
||||||
int width;
|
int width;
|
||||||
TuiCell() : ch(" "), width(1) {}
|
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 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="╯",
|
.top_left="╭", .top_right="╮", .bottom_left="╰", .bottom_right="╯",
|
||||||
.bottom="─", .top="─", .horizontal="─", .vertical="│", .bg=" ",
|
.bottom="─", .top="─", .horizontal="─", .vertical="│", .bg=" ",
|
||||||
.cross="┼", .tee_up="┴", .tee_down="┬", .tee_left="┤", .tee_right="├"
|
.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="+",
|
.top_left="+", .top_right="+", .bottom_left="+", .bottom_right="+",
|
||||||
.bottom="-", .top="-", .vertical="|", .bg=" ", .cross="+",
|
.bottom="-", .top="-", .horizontal="-", .vertical="|", .bg=" ",
|
||||||
.tee_up="+", .tee_down="+", .tee_left="+", .tee_right="+"
|
.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="┛",
|
.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="▛",
|
.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="▛",
|
.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);
|
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 disable_raw_mode();
|
||||||
static void enable_raw_mode();
|
static void enable_raw_mode();
|
||||||
static void move_cursor(int x, int y);
|
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 draw_border(int w, int h, border_charset charset = border_clip, int x = 1, int y = 1);
|
||||||
static void hide_cursor();
|
static void hide_cursor();
|
||||||
static void show_cursor();
|
static void show_cursor();
|
||||||
void set_background(Color clr);
|
static void set_background(Color clr);
|
||||||
void set_foreground(Color clr);
|
static void set_foreground(Color clr);
|
||||||
static void move_cursor_right(int n);
|
static void move_cursor_right(int n);
|
||||||
static void move_cursor_x(int col);
|
static void move_cursor_x(int col);
|
||||||
static void move_cursor_y(int row);
|
static void move_cursor_y(int row);
|
||||||
static void invert_color();
|
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 draw_button(int x, int y, int w, int h, std::string_view label, border_charset charset = border_solid);
|
||||||
static void invert_color_off();
|
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_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 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 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"
|
#include "./../tui.hpp"
|
||||||
|
|
||||||
namespace ui
|
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 <chrono>
|
||||||
#include <csignal>
|
#include <csignal>
|
||||||
#include <cstring>
|
#include <cstring>
|
||||||
#include <sys/ioctl.h>
|
|
||||||
#include <unistd.h>
|
#include <unistd.h>
|
||||||
#include <sys/select.h>
|
#include <sys/select.h>
|
||||||
#include <iostream>
|
#include <iostream>
|
||||||
@ -13,16 +12,23 @@ namespace ui
|
|||||||
{
|
{
|
||||||
Tui::Tui()
|
Tui::Tui()
|
||||||
{
|
{
|
||||||
get_console_size();
|
init_draw();
|
||||||
int panel_w = (width / 100) * 27;
|
|
||||||
std::cout << "Ширина (колонок): " << width << "\n";
|
enable_raw_mode();
|
||||||
std::cout << "Высота (строк): " << height << "\n";
|
}
|
||||||
|
|
||||||
|
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();
|
clear();
|
||||||
hide_cursor();
|
hide_cursor();
|
||||||
set_foreground(Color::White);
|
set_foreground(Color::White);
|
||||||
apply_color();
|
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);
|
root->children.push_back(nullptr);
|
||||||
|
|
||||||
@ -41,8 +47,6 @@ namespace ui
|
|||||||
|
|
||||||
move_cursor(width - 8, height - 1);
|
move_cursor(width - 8, height - 1);
|
||||||
std::cout << "[enter>" <<std::flush;
|
std::cout << "[enter>" <<std::flush;
|
||||||
|
|
||||||
enable_raw_mode();
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Tui::~Tui()
|
Tui::~Tui()
|
||||||
@ -122,9 +126,23 @@ namespace ui
|
|||||||
}
|
}
|
||||||
i++;
|
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) {
|
else if (c >= 32) {
|
||||||
|
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;
|
chat_input += c;
|
||||||
input_changed = true;
|
input_changed = true;
|
||||||
|
}
|
||||||
i++;
|
i++;
|
||||||
}
|
}
|
||||||
else {
|
else {
|
||||||
@ -143,11 +161,22 @@ namespace ui
|
|||||||
} else {
|
} else {
|
||||||
cursor = !cursor;
|
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;
|
last_blink = now;
|
||||||
|
|
||||||
int input_x = ((width / 100) * 27) + 3;
|
int input_x = ((width / 100) * 27) + 3;
|
||||||
move_cursor(input_x, height - 1);
|
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;
|
return running;
|
||||||
}
|
}
|
||||||
|
|
||||||
void Tui::get_console_size()
|
winsize Tui::get_console_size()
|
||||||
{
|
{
|
||||||
winsize w;
|
winsize w;
|
||||||
|
ioctl(STDOUT_FILENO, TIOCGWINSZ, &w);
|
||||||
if (ioctl(STDOUT_FILENO, TIOCGWINSZ, &w) == 0) {
|
return w;
|
||||||
width = w.ws_col;
|
|
||||||
height = w.ws_row;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
void Tui::apply_color() const
|
void Tui::apply_color()
|
||||||
{
|
{
|
||||||
std::cout << "\033["<< foreground << ";" << background <<"m" << std::flush;
|
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)
|
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)));
|
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);
|
draw_border_to_buffer(grid, w, h, charset);
|
||||||
@ -281,7 +306,7 @@ namespace ui
|
|||||||
std::cout << grid[i][j].ch;
|
std::cout << grid[i][j].ch;
|
||||||
}
|
}
|
||||||
if (i < h - 1) {
|
if (i < h - 1) {
|
||||||
std::cout << "\n";
|
std::cout << "\033[1B\033[" << x << "G" << std::flush;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -321,13 +346,17 @@ 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)
|
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) {
|
if (split->type == SplitDesc::VERTICAL) {
|
||||||
int x = offset_x + split->position;
|
int x = offset_x + split->position;
|
||||||
|
if (!grid[0].empty() && x >= 0 && x < static_cast<int>(grid[0].size())) {
|
||||||
for (int i = 0; i < h; i++) {
|
for (int i = 0; i < h; i++) {
|
||||||
int y = offset_y + 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];
|
TuiCell& cell = grid[y][x];
|
||||||
|
|
||||||
if (i == 0) {
|
if (i == 0) {
|
||||||
if (cell.ch == charset.top) cell.ch = std::string(charset.tee_down);
|
if (cell.ch == charset.top) cell.ch = std::string(charset.tee_down);
|
||||||
} else if (i == h - 1) {
|
} else if (i == h - 1) {
|
||||||
@ -340,20 +369,22 @@ namespace ui
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
} else { // HORIZONTAL
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
int y = offset_y + split->position;
|
int y = offset_y + split->position;
|
||||||
|
if (y >= 0 && y < static_cast<int>(grid.size())) {
|
||||||
for (int j = 0; j < w; j++) {
|
for (int j = 0; j < w; j++) {
|
||||||
int x = offset_x + j;
|
int x = offset_x + j;
|
||||||
|
if (x >= 0 && x < static_cast<int>(grid[y].size())) {
|
||||||
TuiCell& cell = grid[y][x];
|
TuiCell& cell = grid[y][x];
|
||||||
|
|
||||||
if (j == 0) {
|
if (j == 0) {
|
||||||
if (cell.ch == charset.vertical) cell.ch = std::string(charset.tee_right);
|
if (cell.ch == charset.vertical) cell.ch = std::string(charset.tee_right);
|
||||||
} else if (j == w - 1) {
|
} else if (j == w - 1) {
|
||||||
if (cell.ch == charset.vertical) cell.ch = std::string(charset.tee_left);
|
if (cell.ch == charset.vertical) cell.ch = std::string(charset.tee_left);
|
||||||
} else {
|
} else {
|
||||||
if (cell.ch == charset.bg) {
|
if (cell.ch == charset.bg) {
|
||||||
// Используем horizontal вместо bottom для семантической корректности горизонтального разделителя
|
|
||||||
cell.ch = std::string(charset.horizontal);
|
cell.ch = std::string(charset.horizontal);
|
||||||
} else if (cell.ch == charset.vertical) {
|
} else if (cell.ch == charset.vertical) {
|
||||||
cell.ch = std::string(charset.cross);
|
cell.ch = std::string(charset.cross);
|
||||||
@ -362,6 +393,8 @@ namespace ui
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
void Tui::draw_border_to_buffer(std::vector<std::vector<TuiCell>>& grid, int w, int h, border_charset charset)
|
void Tui::draw_border_to_buffer(std::vector<std::vector<TuiCell>>& grid, int w, int h, border_charset charset)
|
||||||
{
|
{
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user