Xenith/main.cpp
2026-09-29 23:08:51 +07:00

340 lines
12 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.h"
#include "xenith/core/generate.h"
#include "xenith/core/train.h"
#include "xenith/core/gradcheck.h"
#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 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);
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"));
}
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));
}
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;
if (std::getline(ss, segment, '.')) v.major = std::stoi(segment);
if (std::getline(ss, segment, '.')) v.minor = std::stoi(segment);
if (std::getline(ss, segment, '.')) v.patch = std::stoi(segment);
return v;
}
static bool available_new_version()
{
std::string cmd = "curl -s https://git.bipfr.ru/BIPfR/Xenith/raw/branch/master/version.txt -o ./tmp/remote_version.txt";
system(cmd.c_str());
std::ifstream tmp_file("./tmp/remote_version.txt");
std::string version = "0.0.0";
if (tmp_file.is_open()) {
std::getline(tmp_file, version);
tmp_file.close();
}
Version remote_version = parse_version(version);
std::ifstream file("./version.txt");
version = "0.0.0";
if (file.is_open()) {
std::getline(file, version);
file.close();
}
Version current_version = parse_version(version);
if (remote_version.major > current_version.major) return true;
if (remote_version.minor > current_version.minor) return true;
if (remote_version.patch > current_version.patch) return true;
return false;
}
static int cmd_update(const Args& a)
{
if (available_new_version())
{
namespace fs = std::filesystem;
std::cout << "Удаляем старые файлы..." << 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 std::filesystem::__cxx11::filesystem_error& e) {
std::cerr << "Ошибка удаления " << name << ": " << e.what() << std::endl;
}
}
std::cout << "Скачиваем новую версию..." << std::endl;
int result = system("git clone https://git.bipfr.ru/BIPfR/Xenith /tmp/xenith_update");
if (result != 0) {
std::cerr << "Ошибка при клонировании репозитория" << std::endl;
return -1;
}
result = system("cp -r /tmp/xenith_update/* .");
if (result != 0) {
std::cerr << "Ошибка при копировании файлов" << std::endl;
return -1;
}
system("rm -rf /tmp/xenith_update");
std::cout << "Обновление завершено!" << std::endl;
} else
{
std::cout << "Уже установлена последняя версия! " << std::endl;
}
return 0;
}
int main(int argc, char** argv) {
if (argc < 2) { printf(HELP_TEXT.c_str()); return 1; }
const std::string cmd = argv[1];
if (cmd == "-h" || cmd == "--help" || cmd == "help") { printf(HELP_TEXT.c_str()); 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);
fprintf(stderr, "неизвестная команда: %s\n\n", cmd.c_str());
printf(HELP_TEXT.c_str());
return 1;
}