diff --git a/README.md b/README.md index df632a5..2438542 100644 --- a/README.md +++ b/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` — diff --git a/bin/Xenith b/bin/Xenith index 29b2739..7b0d252 100755 Binary files a/bin/Xenith and b/bin/Xenith differ diff --git a/help.h b/help.h index 18d80d9..e84f8fd 100644 --- a/help.h +++ b/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 \ No newline at end of file diff --git a/main.cpp b/main.cpp index 06edeba..2b54fcc 100644 --- a/main.cpp +++ b/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 #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; diff --git a/version.txt b/version.txt index 05b19b1..fa3de58 100644 --- a/version.txt +++ b/version.txt @@ -1 +1 @@ -0.0.4 \ No newline at end of file +0.0.5 \ No newline at end of file diff --git a/xenith/converter.cpp b/xenith/converter.cpp index ac94854..c967101 100644 --- a/xenith/converter.cpp +++ b/xenith/converter.cpp @@ -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! " diff --git a/xenith/core/checkpoint.cpp b/xenith/core/checkpoint.cpp index 93eef21..109c422 100644 --- a/xenith/core/checkpoint.cpp +++ b/xenith/core/checkpoint.cpp @@ -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; diff --git a/xenith/core/generate.cpp b/xenith/core/generate.cpp index 78a3067..097afd5 100644 --- a/xenith/core/generate.cpp +++ b/xenith/core/generate.cpp @@ -44,7 +44,7 @@ void rope_apply(float* x, int hd, int pos, const std::vector& cs, void forward_token(const Params& p, KvCache& kv, int token, int pos, const std::vector& cs, const std::vector& 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* out_ids) { + std::vector* 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); diff --git a/xenith/core/generate.hpp b/xenith/core/generate.hpp index 5998c3e..a538b13 100644 --- a/xenith/core/generate.hpp +++ b/xenith/core/generate.hpp @@ -18,7 +18,7 @@ struct GenConfig { std::string generate(const Params& p, const Tokenizer& tok, const std::string& prompt, const GenConfig& gc, - std::vector* out_ids = nullptr); + std::vector* out_ids = nullptr, bool fast = false); std::vector generate_batch(const Params& p, const Tokenizer& tok, const std::vector& prompts, diff --git a/xenith/core/gradcheck.cpp b/xenith/core/gradcheck.cpp index f522165..1220c60 100644 --- a/xenith/core/gradcheck.cpp +++ b/xenith/core/gradcheck.cpp @@ -1,6 +1,6 @@ // gradcheck.cpp — проверяем, что обратный проход совпадает с численным // и заодно меряем скорость. -#include "gradcheck.h" +#include "gradcheck.hpp" #include "generate.hpp" #include "train.hpp" #include @@ -14,9 +14,9 @@ namespace xt { namespace { // loss на фиксированном батче, без обратного прохода -float eval_loss(const Params& p, const std::vector& x, const std::vector& y, int n) { +float eval_loss(const Params& p, const std::vector& x, const std::vector& 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(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(); diff --git a/xenith/core/gradcheck.hpp b/xenith/core/gradcheck.hpp index fdcb2d6..dd0f1fb 100644 --- a/xenith/core/gradcheck.hpp +++ b/xenith/core/gradcheck.hpp @@ -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); } diff --git a/xenith/core/model.cpp b/xenith/core/model.cpp index 64b0529..3365098 100644 --- a/xenith/core/model.cpp +++ b/xenith/core/model.cpp @@ -204,7 +204,7 @@ void rope_tables(int block, int head_dim, int base, std::vector& 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); diff --git a/xenith/core/model.hpp b/xenith/core/model.hpp index 7c00419..9b95946 100644 --- a/xenith/core/model.hpp +++ b/xenith/core/model.hpp @@ -81,10 +81,10 @@ struct ForwardCache { void rope_tables(int block, int head_dim, int base, std::vector& cs, std::vector& 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); } \ No newline at end of file diff --git a/xenith/core/tensor.hpp b/xenith/core/tensor.hpp index 129bf73..4fb7f2e 100644 --- a/xenith/core/tensor.hpp +++ b/xenith/core/tensor.hpp @@ -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(); diff --git a/xenith/core/train.hpp b/xenith/core/train.hpp index 13da145..db714b6 100644 --- a/xenith/core/train.hpp +++ b/xenith/core/train.hpp @@ -29,6 +29,7 @@ struct TrainConfig { int ckpt_every = 0; int val_every = 0; int val_tokens = 20000; + bool fast = false; }; struct Dataset { diff --git a/xenith/core/trainer.cpp b/xenith/core/trainer.cpp index ac6a3f1..ede0c2c 100644 --- a/xenith/core/trainer.cpp +++ b/xenith/core/trainer.cpp @@ -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 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); diff --git a/xenith/ui/tui.hpp b/xenith/ui/tui.hpp index f1a865b..6c520f9 100644 --- a/xenith/ui/tui.hpp +++ b/xenith/ui/tui.hpp @@ -1,8 +1,11 @@ #pragma once +#include #include #include +#include #include +#include #include #include @@ -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 left = std::nullopt; + std::optional right = std::nullopt; + }; + inline static std::vector chat_history; struct SplitDesc { enum Type { VERTICAL, HORIZONTAL }; Type type; int position; std::vector> 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>& grid, int offset_x, int offset_y, int w, int h, border_charset charset, const SplitDesc* split); static void apply_single_split(std::vector>& 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>& 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(); } \ No newline at end of file diff --git a/xenith/ui/tui/message.cpp b/xenith/ui/tui/message.cpp index 78bd102..fc04f2c 100644 --- a/xenith/ui/tui/message.cpp +++ b/xenith/ui/tui/message.cpp @@ -1,7 +1,68 @@ -#include "message.h" +#include +#include + #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); + } + } -} \ No newline at end of file + 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::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(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(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; + } + + } +} diff --git a/xenith/ui/tui/tui.cpp b/xenith/ui/tui/tui.cpp index ed95687..9d4b4b0 100644 --- a/xenith/ui/tui/tui.cpp +++ b/xenith/ui/tui/tui.cpp @@ -3,7 +3,6 @@ #include #include #include -#include #include #include #include @@ -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::VERTICAL, (width / 100) * 27); + const auto root = std::make_unique(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>" <= 32) { - chat_input += c; - input_changed = true; + int spaces_to_draw = std::max(0, (width - ((width / 100) * 27) - 15 - static_cast(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(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> grid(h, std::vector(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>& 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(grid[0].size())) { + for (int i = 0; i < h; i++) { + int y = offset_y + i; + if (y >= 0 && y < static_cast(grid.size())) { + if (x < static_cast(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(grid.size())) { + for (int j = 0; j < w; j++) { + int x = offset_x + j; + if (x >= 0 && x < static_cast(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); + } + } } } }