Xenith/main.cpp
2026-10-04 13:57:10 +07:00

436 lines
16 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#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 <cstdio>
#include "help.h"
#include <cstdlib>
#include <filesystem>
#include <string>
#include <vector>
#include <map>
#include <fstream>
#include <iostream>
#include <sstream>
#include <sys/stat.h>
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<std::string, std::string> kv;
std::vector<std::string> 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<std::string> 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<std::string> 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<int> 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);
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);
}
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 false;
}
std::cout << "\033[32mДоступна новая версия \033[33m" << remote_str << "\033[32m\nВыполните \"Xenith update\" чтобы обновится до последней версии\033[39m\n" << 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;
}