BigMoeOnEdge/cli/main.cpp

159 lines
7.3 KiB
C++

// bmoe-cli — host driver for BigMoeOnEdge.
//
// Parses flags into a RunConfig and runs the engine. Two output modes:
// * default: streams the generated text inline, per-token timing to stderr;
// * --progress: one machine-readable JSON line per token (docs/telemetry.md), which
// the Android example app parses for its live panel.
//
// Environment variables are read ONLY here, as overrides for the matching flags, so the
// engine stays env-free. The flag always wins over the env value.
#include "bmoe/config.h"
#include "bmoe/runtime.h"
#include "bmoe/recipe.h"
#include "bmoe/metrics.h"
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <memory>
#include <string>
using namespace bmoe;
static int env_int(const char * k, int dflt) {
const char * v = std::getenv(k);
return (v && *v) ? std::atoi(v) : dflt;
}
static std::string json_escape(const std::string & s) {
std::string o;
o.reserve(s.size() + 8);
for (char c : s) {
switch (c) {
case '"': o += "\\\""; break;
case '\\': o += "\\\\"; break;
case '\n': o += "\\n"; break;
case '\r': o += "\\r"; break;
case '\t': o += "\\t"; break;
default:
if ((unsigned char) c < 0x20) { char b[8]; std::snprintf(b, sizeof(b), "\\u%04x", c); o += b; }
else o += c;
}
}
return o;
}
static void print_usage(const char * argv0) {
std::printf(
"usage: %s -m <model.gguf> [options]\n"
"\n"
" -m, --model PATH gguf model (required)\n"
" -p, --prompt STR prompt text\n"
" -n, --n-predict N tokens to generate (default 32)\n"
" -t, --threads N compute threads (default 4)\n"
" -c, --ctx-size N context size (default 2048)\n"
" --chatml wrap the prompt in a Qwen ChatML turn\n"
" --progress emit machine telemetry (one JSON line per token)\n"
" --csv PATH also write per-token metrics as CSV\n"
"\n"
" MoE expert streaming:\n"
" --moe-stream stream only the routed experts per token (MoE models)\n"
" --cache-mb N LRU expert cache budget in MiB (0=off, or >=%d)\n"
" --io-threads N parallel expert-read lanes [1..%d] (default 4)\n"
" --no-odirect do not bypass the page cache\n"
" --load-all debug: read ALL experts each token (A/B baseline)\n"
" --force-cache allow a cache-mb in the pathological band\n"
" --list-archs print supported MoE architectures and exit\n"
"\n"
" Env overrides (flag wins): BMOE_CACHE_MB, BMOE_IO_THREADS, BMOE_PROGRESS\n",
argv0, MoeStreamConfig::cache_min_mb, MoeStreamConfig::io_threads_max);
}
int main(int argc, char ** argv) {
RunConfig cfg;
std::string csv_path;
for (int i = 1; i < argc; ++i) {
std::string a = argv[i];
auto next = [&](const char * what) -> const char * {
if (i + 1 >= argc) { std::fprintf(stderr, "missing value for %s\n", what); std::exit(1); }
return argv[++i];
};
if (a == "-m" || a == "--model") cfg.model_path = next("-m");
else if (a == "-p" || a == "--prompt") cfg.prompt = next("-p");
else if (a == "-n" || a == "--n-predict") cfg.n_predict = std::atoi(next("-n"));
else if (a == "-t" || a == "--threads") cfg.n_threads = std::atoi(next("-t"));
else if (a == "-c" || a == "--ctx-size") cfg.n_ctx = std::atoi(next("-c"));
else if (a == "--chatml") cfg.chatml = true;
else if (a == "--progress") cfg.progress = true;
else if (a == "--csv") csv_path = next("--csv");
else if (a == "--moe-stream") cfg.moe.enabled = true;
else if (a == "--cache-mb") cfg.moe.cache_mb = std::atoi(next("--cache-mb"));
else if (a == "--io-threads") cfg.moe.io_threads = std::atoi(next("--io-threads"));
else if (a == "--no-odirect") cfg.moe.o_direct = false;
else if (a == "--load-all") cfg.moe.load_all = true;
else if (a == "--force-cache") cfg.moe.force_cache = true;
else if (a == "--list-archs") {
std::printf("supported MoE architectures:\n");
for (int k = 0; k < n_moe_recipes(); ++k) std::printf(" %s\n", moe_recipe_at(k)->arch);
return 0;
}
else if (a == "-h" || a == "--help") { print_usage(argv[0]); return 0; }
else { std::fprintf(stderr, "unknown arg: %s\n", a.c_str()); print_usage(argv[0]); return 1; }
}
// Env overrides (flag wins: only apply when the flag left the default).
if (cfg.moe.cache_mb == 0) cfg.moe.cache_mb = env_int("BMOE_CACHE_MB", 0);
if (cfg.moe.io_threads == 4) cfg.moe.io_threads = env_int("BMOE_IO_THREADS", 4);
if (!cfg.progress) cfg.progress = env_int("BMOE_PROGRESS", 0) != 0;
if (cfg.model_path.empty()) { print_usage(argv[0]); return 1; }
ValidationResult vr = validate(cfg);
if (!vr) { std::fprintf(stderr, "config error: %s\n", vr.error.c_str()); return 1; }
std::unique_ptr<IMetricsSink> sink;
if (!csv_path.empty()) {
sink.reset(make_csv_metrics_sink(csv_path));
if (!sink) std::fprintf(stderr, "warning: could not open csv %s\n", csv_path.c_str());
}
if (!cfg.progress) { std::printf("%s", cfg.prompt.c_str()); std::fflush(stdout); }
auto on_token = [&](const TokenMetrics & m) {
if (cfg.progress) {
if (m.read_bytes || m.io_ms > 0.0)
std::printf("BMOE_LOAD {\"mb\":%.2f,\"ms\":%.1f}\n", m.read_bytes / (1024.0 * 1024.0), m.io_ms);
std::printf("BMOE_PROGRESS {\"step\":%d,\"steps\":%d,\"wall_ms\":%.1f,\"io_ms\":%.1f,"
"\"compute_ms\":%.1f,\"cache_hit_pct\":%.1f,\"text\":\"%s\"}\n",
m.step, m.steps, m.wall_ms, m.io_ms, m.compute_ms, m.cache_hit_pct,
json_escape(m.text).c_str());
std::fflush(stdout);
} else {
std::fwrite(m.piece.data(), 1, m.piece.size(), stdout);
std::fflush(stdout);
}
};
RunResult r = run(cfg, on_token, sink.get());
if (!r) { std::fprintf(stderr, "\nerror: %s\n", r.error.c_str()); return 1; }
const RunSummary & s = r.summary;
if (cfg.progress) {
std::printf("=== answer ===\n%s\n=== perf ===\n", r.generated_text.c_str());
} else {
std::printf("\n\n");
}
std::printf("generation: %d tokens, %.3f s/token (%.3f tok/s)\n",
s.n_generated, s.s_per_token, s.tokens_per_second);
if (cfg.moe.enabled) {
std::printf("moe-stream: read %.1f MiB (%.2f MiB/token), decode %.3f s/token "
"(compute %.3f + flash I/O %.3f s/token, %.0f MiB/s)\n",
s.moe_read_mib, s.n_generated ? s.moe_read_mib / s.n_generated : 0.0,
s.s_per_token, s.moe_compute_s_per_token, s.moe_io_s_per_token,
s.moe_io_seconds > 0 ? s.moe_read_mib / s.moe_io_seconds : 0.0);
if (s.cache_hit_pct >= 0.0)
std::printf("moe-cache: %.1f%% hit, resident %.1f MiB\n", s.cache_hit_pct, s.cache_resident_mib);
}
return 0;
}