mirror of
https://github.com/Helldez/BigMoeOnEdge.git
synced 2026-10-03 19:45:46 +00:00
159 lines
7.3 KiB
C++
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;
|
|
}
|