Xenith/main.cpp
2026-09-29 20:21:46 +07:00

370 lines
15 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.

// main.cpp — командная строка Xenith
#include "xenith/core/checkpoint.h"
#include "xenith/core/generate.h"
#include "xenith/core/train.h"
#include "xenith/core/gradcheck.h"
#include <cstdio>
#include <cstring>
#include <cstdlib>
#include <string>
#include <vector>
#include <map>
#include <fstream>
#include <sstream>
#include <sys/stat.h>
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<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 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<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;
}
// ------------------------------------------------------------------ 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<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;
}
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;
}