#include "xenith/core/checkpoint.hpp" #include "xenith/core/generate.hpp" #include "xenith/core/train.hpp" #include "xenith/core/gradcheck.hpp" #include "xenith/ui/tui.hpp" #include #include "help.h" #include #include #include #include #include #include #include #include #include namespace fs = std::filesystem; namespace std::filesystem::__cxx11 { class filesystem_error; } using namespace xt; static std::string read_all(const std::string& p, bool& ok) { std::ifstream f(p, std::ios::binary); if (!f) { ok = false; return {}; } std::ostringstream ss; ss << f.rdbuf(); ok = true; return ss.str(); } static void mkdirs(const std::string& path) { std::string cur; for (size_t i = 0; i <= path.size(); ++i) { if (i == path.size() || path[i] == '/') { if (!cur.empty() && cur != ".") ::mkdir(cur.c_str(), 0755); if (i == path.size()) break; } cur += path[i]; } } struct Version { int major = 0; int minor = 0; int patch = 0; }; struct Args { std::map kv; std::vector pos; Args(int argc, char** argv, int from) { for (int i = from; i < argc; ++i) { std::string a = argv[i]; if (a.rfind("--", 0) == 0) { std::string key = a.substr(2); size_t eq = key.find('='); if (eq != std::string::npos) { kv[key.substr(0, eq)] = key.substr(eq + 1); } else if (i + 1 < argc && std::string(argv[i + 1]).rfind("--", 0) != 0) { kv[key] = argv[++i]; } else { kv[key] = "1"; } } else { pos.push_back(a); } } } bool has(const char* k) const { return kv.count(k) > 0; } std::string get(const char* k, const std::string& def = "") const { auto it = kv.find(k); return it == kv.end() ? def : it->second; } int geti(const char* k, int def) const { auto it = kv.find(k); return it == kv.end() ? def : std::atoi(it->second.c_str()); } float getf(const char* k, float def) const { auto it = kv.find(k); return it == kv.end() ? def : (float)std::atof(it->second.c_str()); } }; static int cmd_new(const Args& a) { if (a.pos.empty()) { fprintf(stderr, "new: укажите путь модели\n"); return 1; } const std::string out = a.pos[0]; const std::string corpus = a.get("corpus"); if (corpus.empty()) { fprintf(stderr, "new: нужен --corpus\n"); return 1; } bool ok = false; std::string text = read_all(corpus, ok); if (!ok) { fprintf(stderr, "new: не читается %s\n", corpus.c_str()); return 1; } Checkpoint ck; ck.seed = (uint64_t)a.geti("seed", 1337); ck.params.cfg.vocab_size = a.geti("vocab", 512); ck.params.cfg.n_embd = a.geti("embd", 128); ck.params.cfg.n_layer = a.geti("layers", 4); ck.params.cfg.n_head = a.geti("heads", 4); ck.params.cfg.block_size = a.geti("ctx", 128); ck.params.cfg.ffn_dim = a.geti("ffn", 0); ck.params.cfg.rope_base = a.geti("rope", 10000); if (a.has("untied")) ck.params.cfg.n_vocab_out = ck.params.cfg.vocab_size; ck.tok.build_from_text(text, ck.params.cfg.vocab_size, a.has("view-vocab")); ck.params.cfg.vocab_size = ck.tok.size(); std::string err; if (!ck.params.cfg.valid(err)) { fprintf(stderr, "new: плохая конфигурация: %s\n", err.c_str()); return 1; } if (ck.tok.size() < 16) { fprintf(stderr, "new: корпус дал всего %d токенов, нужно больше текста\n", ck.tok.size()); return 1; } ck.params.init(ck.seed); printf("создаю модель: %s\n", out.c_str()); print_model_info(ck); size_t slash = out.find_last_of('/'); if (slash != std::string::npos) mkdirs(out.substr(0, slash)); if (!save_checkpoint(out, ck, err)) { fprintf(stderr, "new: %s\n", err.c_str()); return 1; } printf("готово. обучать: xenith train %s --corpus %s --out trained.xnh\n", out.c_str(), corpus.c_str()); return 0; } static int cmd_train(const Args& a) { if (a.pos.empty()) { fprintf(stderr, "train: укажите файл модели\n"); return 1; } const std::string path = a.pos[0]; const std::string corpus = a.get("corpus"); if (corpus.empty()) { fprintf(stderr, "train: нужен --corpus\n"); return 1; } Checkpoint ck; std::string err; if (!load_checkpoint(path, ck, err)) { fprintf(stderr, "train: %s\n", err.c_str()); return 1; } TrainConfig tc; tc.steps = a.geti("steps", 1000); tc.batch_size = a.geti("batch", 8); tc.block = a.geti("block", 64); tc.lr = a.getf("lr", 3e-4f); tc.weight_decay = a.getf("wd", 0.01f); tc.clip = a.getf("clip", 1.0f); tc.warmup = a.geti("warmup", 100); tc.seed = (uint64_t)a.geti("seed", 1337); tc.log_every = a.geti("log-every", 50); tc.ckpt_every = a.geti("ckpt-every", 0); 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; printf("обучаю %s\n", path.c_str()); print_model_info(ck); printf(" шагов %d, батч %d, окно %d, lr %g\n\n", tc.steps, tc.batch_size, tc.block, tc.lr); int resume = a.has("resume") ? ck.step : 0; if (resume > 0) { fprintf(stderr, "продолжаю с шага %d\n", resume); } std::string out = a.get("out", path); train_model(ck.params, ck.tok, corpus, tc, out, resume); return 0; } static int cmd_gen(const Args& a) { if (a.pos.empty()) { fprintf(stderr, "gen: укажите файл модели\n"); return 1; } Checkpoint ck; std::string err; if (!load_checkpoint(a.pos[0], ck, err)) { fprintf(stderr, "gen: %s\n", err.c_str()); return 1; } GenConfig gc; gc.max_new_tokens = a.geti("n", 200); gc.temperature = a.getf("temp", 0.8f); gc.top_k = a.geti("top-k", 40); gc.top_p = a.getf("top-p", 0.95f); gc.seed = a.geti("seed", -1); gc.stream = !a.has("no-stream"); if (a.has("batch")) { std::vector prompts; std::string s = a.get("batch"); size_t start = 0; while (start <= s.size()) { size_t p = s.find(';', start); if (p == std::string::npos) { prompts.push_back(s.substr(start)); break; } prompts.push_back(s.substr(start, p - start)); start = p + 1; } gc.stream = false; std::vector outs = generate_batch(ck.params, ck.tok, prompts, gc); for (size_t i = 0; i < outs.size(); ++i) { printf("=== промпт %zu: %s\n%s\n\n", i + 1, prompts[i].c_str(), outs[i].c_str()); } return 0; } std::string prompt = a.get("prompt", ""); if (a.pos.size() > 1) prompt = a.pos[1]; std::vector ids; std::string text = generate(ck.params, ck.tok, prompt, gc, &ids); if (a.has("show-tokens")) { printf("\n--- токены (%zu) ---\n", ids.size()); for (size_t i = 0; i < ids.size(); ++i) printf("%d ", ids[i]); printf("\n"); } return 0; } static int cmd_info(const Args& a) { if (a.pos.empty()) { fprintf(stderr, "info: укажите файл модели\n"); return 1; } Checkpoint ck; std::string err; if (!load_checkpoint(a.pos[0], ck, err)) { fprintf(stderr, "info: %s\n", err.c_str()); return 1; } printf("файл: %s\n", a.pos[0].c_str()); print_model_info(ck); printf(" словарь (первые 40): "); for (int i = 0; i < 40 && i < ck.tok.size(); ++i) printf("%d=%s ", i, ck.tok.id2tok[i].c_str()); printf("%s\n", ck.tok.size() > 40 ? " ..." : ""); return 0; } static int cmd_gradcheck(const Args& a) { return run_gradcheck(a.has("verbose"), a.has("fast")); } static int cmd_bench(const Args& a) { if (a.pos.empty()) { fprintf(stderr, "bench: укажите файл модели\n"); return 1; } 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), a.has("fast")); } static int cmd_ollama(const Args& a) { fprintf(stderr, "На стадии разработки"); return 0; } Version parse_version(const std::string& str) { Version v; std::stringstream ss(str); std::string segment; try { if (std::getline(ss, segment, '.')) v.major = std::stoi(segment); else return v; if (std::getline(ss, segment, '.')) v.minor = std::stoi(segment); else return v; if (std::getline(ss, segment, '.')) v.patch = std::stoi(segment); } catch (const std::exception& e) { std::cerr << "Ошибка парсинга версии '" << str << "': " << e.what() << std::endl; } return v; } bool is_newer(const Version& remote, const Version& current) { if (remote.major != current.major) return remote.major > current.major; if (remote.minor != current.minor) return remote.minor > current.minor; return remote.patch > current.patch; } static bool available_new_version() { const std::string REMOTE_URL = "https://git.bipfr.ru/BIPfR/Xenith/raw/branch/master/version.txt"; const std::string TMP_FILE = "./tmp/remote_version.txt"; fs::create_directories("./tmp"); std::string cmd = "curl -s --fail \"" + REMOTE_URL + "\" -o \"" + TMP_FILE + "\""; int curl_result = system(cmd.c_str()); if (curl_result != 0) { std::cerr << "Не удалось скачать версию. Проверьте интернет или URL." << std::endl; return false; } std::ifstream tmp_file(TMP_FILE); std::string remote_str = "0.0.0"; if (tmp_file.is_open() && std::getline(tmp_file, remote_str)) { tmp_file.close(); } else { std::cerr << "Файл версии пуст или нечитаем" << std::endl; return false; } std::ifstream local_file("./version.txt"); std::string local_str = "0.0.0"; if (local_file.is_open() && std::getline(local_file, local_str)) { local_file.close(); } Version remote_ver = parse_version(remote_str); Version local_ver = parse_version(local_str); return is_newer(remote_ver, local_ver); } static int cmd_update(const Args& a) { if (!available_new_version()) { std::cout << "Уже установлена последняя версия!" << std::endl; return 0; } std::cout << "Начинаем обновление..." << std::endl; std::cout << "Очистка старых файлов (сохраняем models/)..." << std::endl; for (const auto& entry : fs::directory_iterator(".")) { std::string name = entry.path().filename().string(); if (name == "models") continue; try { if (fs::is_directory(entry.status())) { fs::remove_all(entry.path()); } else { fs::remove(entry.path()); } } catch (const fs::filesystem_error& e) { std::cerr << "Ошибка удаления " << name << ": " << e.what() << std::endl; } } const std::string REPO_URL = "https://git.bipfr.ru/BIPfR/Xenith.git"; const std::string TMP_CLONE = "/tmp/xenith_update"; 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()); if (result != 0) { std::cerr << "Ошибка клонирования репозитория" << std::endl; return -1; } std::cout << "Копирование файлов..." << std::endl; result = system(("cp -ra " + TMP_CLONE + "/. .").c_str()); if (result != 0) { std::cerr << "Ошибка копирования файлов" << std::endl; 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; } 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) { auto app = ui::Tui(); while (app.is_running()) { app.update(); } std::cout << "\033[2J\033[H" << std::flush; return 0; } int main(int argc, char** argv) { if (available_new_version()) { const std::string TMP_FILE = "./tmp/remote_version.txt"; std::ifstream tmp_file(TMP_FILE); std::string remote_str = "0.0.0"; if (tmp_file.is_open() && std::getline(tmp_file, remote_str)) { tmp_file.close(); } else { std::cerr << "Ошибка чтения файла новой версии (Пожалуйста напишите о этом баге)" << std::endl; return 1; } std::cout << "\033[32mДоступна новая версия \033[33m" << remote_str << "\033[32m\nВыполните \"Xenith update\" чтобы обновится до последней версии\033[39m" << std::endl; } if (argc < 2) { printf(HELP_TEXT); return 1; } const std::string cmd = argv[1]; if (cmd == "-h" || cmd == "--help" || cmd == "help") { printf(HELP_TEXT); return 0; } Args a(argc, argv, 2); if (cmd == "new") return cmd_new(a); if (cmd == "train") return cmd_train(a); if (cmd == "gen") return cmd_gen(a); if (cmd == "info") return cmd_info(a); if (cmd == "gradcheck") return cmd_gradcheck(a); if (cmd == "bench") return cmd_bench(a); if (cmd == "ollama") return cmd_ollama(a); if (cmd == "update") return cmd_update(a); if (cmd == "tui") return cmd_tui(a); fprintf(stderr, "неизвестная команда: %s\n\n", cmd.c_str()); printf(HELP_TEXT); return 1; }