388 lines
14 KiB
C++
388 lines
14 KiB
C++
#include "xenith/core/checkpoint.h"
|
||
#include "xenith/core/generate.h"
|
||
#include "xenith/core/train.h"
|
||
#include "xenith/core/gradcheck.h"
|
||
#include "xenith/ui/tui.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 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);
|
||
|
||
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;
|
||
|
||
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);
|
||
|
||
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";
|
||
|
||
system(("rm -rf " + TMP_CLONE).c_str());
|
||
|
||
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;
|
||
system(("rm -rf " + TMP_CLONE).c_str());
|
||
return -1;
|
||
}
|
||
system(("rm -rf " + TMP_CLONE).c_str());
|
||
std::cout << "Обновление завершено!" << std::endl;
|
||
return 0;
|
||
}
|
||
|
||
static int cmd_tui(const Args& a)
|
||
{
|
||
ui::Tui tui = ui::Tui();
|
||
while (tui.is_running())
|
||
{
|
||
tui.update();
|
||
}
|
||
return 0;
|
||
}
|
||
|
||
int main(int argc, char** argv) {
|
||
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;
|
||
} |