// main.cpp — командная строка Xenith #include "xenith/core/checkpoint.h" #include "xenith/core/generate.h" #include "xenith/core/train.h" #include "xenith/core/gradcheck.h" #include #include #include #include #include #include #include #include #include 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]; } } // Разбор аргументов вида --key value 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 void usage() { printf(R"(ИСПОЛЬЗОВАНИЕ xenith <команда> [опции] КОМАНДЫ new создать новую модель из корпуса train обучить существующую модель gen сгенерировать текст info показать конфигурацию модели gradcheck проверить градиенты численно bench замерить скорость форварда и шага обучения ollama запустьтить сервер с api ollama ПРИМЕРЫ xenith new model.xnh --corpus data.txt --vocab 512 --embd 64 --layers 4 --heads 4 xenith train model.xnh --corpus data.txt --out model_trained.xnh --steps 3000 xenith train model_trained.xnh --corpus data.txt --out model_trained.xnh --steps 6000 xenith gen model_trained.xnh --prompt "Привет" --n 200 --temp 0.8 ОПЦИИ new --corpus PATH текст для построения словаря (обязательно) --out PATH файл модели (по умолчанию model.xnh) --vocab N размер словаря (по умолчанию 512) --embd N размер эмбеддинга (по умолчанию 128) --layers N число слоёв (по умолчанию 4) --heads N число голов (по умолчанию 4) --ctx N максимальный контекст (по умолчанию 128) --ffn N размер FFN (по умолчанию 4*embd) --rope N база RoPE (по умолчанию 10000) --seed N зерно инициализации (по умолчанию 1337) --untied отдельная матрица выхода вместо привязанной к wte ОПЦИИ train --corpus PATH обучающий текст (обязательно) --out PATH куда сохранить (по умолчанию перезапись входного файла) --steps N шагов (по умолчанию 1000) --batch N батч (по умолчанию 8) --block N длина окна (по умолчанию 64) --lr F скорость (по умолчанию 3e-4) --wd F weight decay (по умолчанию 0.01) --clip F клип нормы градиента (по умолчанию 1.0) --warmup N прогрев (по умолчанию 100) --seed N зерно (по умолчанию 1337) --log-every N как часто печатать (по умолчанию 50) --ckpt-every N промежуточное сохранение каждые N шагов (0 = нет) --val-every N валидация каждые N шагов (0 = нет) --val-tokens N размер валидации (по умолчанию 20000) --threads N потоков (0 = все ядра) --resume продолжить с сохранённого step (Adam-state из файла) ОПЦИИ gen --prompt STR стартовый текст --n N сколько токенов (по умолчанию 200) --temp F температура, 0 = жадный выбор (по умолчанию 0.8) --top-k N top-k (0 = выкл) --top-p F top-p (1.0 = выкл) --seed N зерно --no-stream не печатать в процессе --batch "a;b;c" несколько промптов через ; --show-tokens показать id токенов ОПЦИИ ollama --port PORT порт на котором будет открыт сервер --conf NAME.conf )"); } // ------------------------------------------------------------------ new 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); 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); // размер папки, если out = models/xxx.xnh 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; } // ------------------------------------------------------------------ train 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) { // Adam-state уже загружен из файла, но steps надо расширить fprintf(stderr, "продолжаю с шага %d\n", resume); } std::string out = a.get("out", path); train_model(ck.params, ck.tok, corpus, tc, out, resume); return 0; } // ------------------------------------------------------------------ gen 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; } // ------------------------------------------------------------------ info 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) { 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; } int main(int argc, char** argv) { if (argc < 2) { usage(); return 1; } const std::string cmd = argv[1]; if (cmd == "-h" || cmd == "--help" || cmd == "help") { usage(); 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_bench(a); fprintf(stderr, "неизвестная команда: %s\n\n", cmd.c_str()); usage(); return 1; }