Files
hot-step-cpp-ROCm/engine/tools/hot-step-server.cpp
T
2026-08-16 18:24:52 +07:00

4397 lines
190 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.
// hot-step-server.cpp: HOT-Step HTTP server for ACE-Step music generation
//
// Based on upstream ace-server.cpp with HOT-Step extensions:
// - VAE model selection (multiple VAEs via vae_model field)
// - /vram endpoint for GPU memory reporting
// - Output format from URL ?format= param (backward compat)
// - Adapter absolute-path fallback
//
// Single binary, one port. All compute endpoints (POST /lm, POST /synth,
// POST /understand) are asynchronous: they validate the request, create a
// job, push it to a FIFO queue, and return the job ID immediately.
// A single worker thread processes jobs in order.
// Clients poll GET /job?id=N for status and fetch results with
// GET /job?id=N&result=1. POST /job?id=N&cancel=1 cancels a job.
//
// Job IDs are random 64-bit hex strings (non-predictable).
// Completed jobs are evicted FIFO when the pool exceeds MAX_JOBS.
// Running jobs are never evicted.
//
// Models are discovered by scanning --models directory at startup
// (reads GGUF metadata only, no weights loaded).
// Each request loads the model, executes, and frees it. No model persists
// in VRAM between requests unless --keep-loaded is set. GPU access is
// serialized by the single worker thread (no mutex needed).
//
// Available models are classified by their GGUF general.architecture:
// acestep-lm -> lm bucket
// acestep-dit -> dit bucket
// acestep-text-enc -> text-enc bucket (singleton, first entry used)
// acestep-vae -> vae bucket (singleton, first entry used)
//
// Endpoint requirements:
// /lm LM
// /synth DiT + Text-Enc + VAE
// /understand LM + DiT + VAE
#include "adapter-cancel.h" // g_adapter_cancel — set by worker around ace_synth_load
#include "audio-io.h"
#include "audio-resample.h"
#include "denoiser.h"
#include "spectral-lifter.h"
#include "supersep.h"
#include "hot-step-params.h"
#include "lua-plugin-registry.h"
// ── Linker guard: verify hot-step-sampler.h is active ────────────────
// hot-step-sampler.h defines hotstep_sampler_linked_ with external linkage.
// pipeline-synth-ops.cpp includes it, compiling the symbol into acestep-core.lib.
// If upstream sync clobbers the include back to dit-sampler.h, this symbol
// vanishes and the linker fails here — making the regression a build error.
extern int hotstep_sampler_linked_;
static volatile int * _hotstep_guard_ = &hotstep_sampler_linked_;
#include "model-registry.h"
#include "model-store.h"
#include "vae.h"
#include "vae-enc.h"
#include "pipeline-lm.h"
#include "pipeline-synth.h"
#include "pipeline-understand.h"
#include "request.h"
#include "synth-batch-runner.h"
#include "task-types.h"
#include "timer.h"
#include "version.h"
#include "yyjson.h"
// embedded webui (generated by xxd.cmake from tools/webui/public/index.html.gz)
#include "index.html.gz.hpp"
// suppress warnings in third-party headers
#ifdef __GNUC__
# pragma GCC diagnostic push
# pragma GCC diagnostic ignored "-Wshadow"
#endif
#include "httplib.h"
#ifdef __GNUC__
# pragma GCC diagnostic pop
#endif
#include <atomic>
#include <condition_variable>
#include <csignal>
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <deque>
#include <functional>
#include <memory>
#include <mutex>
#include <random>
#include <string>
#include <thread>
#include <unordered_map>
#include <vector>
#ifdef _WIN32
# include <fcntl.h>
# include <io.h>
# ifndef STDERR_FILENO
# define STDERR_FILENO 2
# endif
#else
# include <unistd.h>
#endif
#ifdef GGML_USE_CUDA
# include "../src/gpu.h"
#endif
// portable fd wrappers. avoids macros that collide with C++ method names
// (e.g. sink.write() in httplib would be eaten by a write() macro).
#ifdef _WIN32
static int fd_pipe(int fd[2]) {
return _pipe(fd, 4096, _O_BINARY);
}
static int fd_dup(int fd) {
return _dup(fd);
}
static int fd_dup2(int src, int dst) {
return _dup2(src, dst);
}
static int fd_read(int fd, void * buf, size_t n) {
return _read(fd, buf, (unsigned) n);
}
static int fd_write(int fd, const void * buf, size_t n) {
return _write(fd, buf, (unsigned) n);
}
static void fd_close(int fd) {
_close(fd);
}
#else
static int fd_pipe(int fd[2]) {
return pipe(fd);
}
static int fd_dup(int fd) {
return dup(fd);
}
static int fd_dup2(int src, int dst) {
return dup2(src, dst);
}
static int fd_read(int fd, void * buf, size_t n) {
return (int) read(fd, buf, n);
}
static int fd_write(int fd, const void * buf, size_t n) {
return (int) write(fd, buf, n);
}
static void fd_close(int fd) {
close(fd);
}
#endif
// server instance pointer for the signal handler
static httplib::Server * g_svr = nullptr;
static void on_signal(int) {
if (g_svr) {
g_svr->stop();
}
}
// work queue: all GPU jobs go through a single FIFO queue processed
// by one worker thread. GPU access is serialized by construction.
static std::deque<std::function<void()>> g_work_queue;
static std::mutex mtx_work;
static std::condition_variable cv_work;
static bool g_work_stop = false;
static void work_push(std::function<void()> fn) {
std::lock_guard<std::mutex> lock(mtx_work);
g_work_queue.push_back(std::move(fn));
cv_work.notify_one();
}
// worker thread: consume jobs in FIFO order until shutdown.
// on stop: finishes the current job, discards pending ones.
static void worker_main() {
for (;;) {
std::function<void()> fn;
{
std::unique_lock<std::mutex> lock(mtx_work);
cv_work.wait(lock, [] { return g_work_stop || !g_work_queue.empty(); });
if (g_work_stop) {
break;
}
fn = std::move(g_work_queue.front());
g_work_queue.pop_front();
}
fn();
}
}
// central GGML module store shared across pipelines. Policy picked at startup
// from --keep-loaded: STRICT by default (one GPU module resident at a time),
// NEVER when the flag is set (accumulate across requests).
static ModelStore * g_store = nullptr;
// model registry (populated at startup from GGUF metadata)
static ModelRegistry g_registry;
// loaded model names (empty = nothing loaded)
static std::string g_loaded_lm;
static std::string g_loaded_dit;
static std::string g_loaded_adapter;
static float g_loaded_adapter_scale = 1.0f;
static std::string g_loaded_und_dit;
static std::string g_loaded_vae;
// pipeline params (rebuilt from registry paths on each load)
static AceLmParams g_lm_params;
static AceSynthParams g_synth_params;
static AceUnderstandParams g_und_params;
// limits
static int g_max_batch = 1;
static int g_mp3_kbps = 128;
static bool g_keep_loaded = false;
// True only when --keep-loaded was on the command line — that co-residency is
// the user's explicit choice and /models/restore-policy must refuse to undo
// it. A ?keep_loaded=1 request latch, by contrast, is transient job plumbing
// (the codes audition) and IS restorable.
static bool g_keep_loaded_cli = false;
// speculative decoding: path to 0.6B draft model (auto-discovered or --draft-lm)
static std::string g_draft_lm_path;
// ONNX model directory (optional, for TensorRT/CUDA EP accelerated VAE)
static const char * g_onnx_dir = nullptr;
// HOT-Step: pre-computed noise profile for spectral denoiser.
// Loaded once at startup from a reference noise sample WAV.
static NoiseProfile g_noise_profile;
// latent format constants (matching upstream ace-server.cpp)
static const int MAX_T_LATENT = 15000; // ~10min at 25Hz
static const int LATENT_CHANNELS = 64;
static const int LATENT_FRAME_BYTES = LATENT_CHANNELS * (int) sizeof(float);
// job system: all compute endpoints create a job and return its ID
// immediately. the worker thread processes jobs in FIFO order, stores
// the result. the client polls GET /job?id=N until done, then fetches
// the result with GET /job?id=N&result=1.
// cancel: POST /job?id=N&cancel=1 sets the per-job flag.
// Fine-grained phase the worker is currently in. Surfaced through GET /job
// (and GET /jobs) alongside the coarse int status so the wrapper can tell
// *why* a long-running job is taking a while — model load vs. adapter
// precompute (the ~17 s LoKr stall) vs. actual DiT inference. Phases are
// advisory: workers may skip ones that don't apply (e.g. /lm never enters
// real DIT_INFERENCE). Order roughly follows the synth pipeline.
enum class JobPhase : int {
QUEUED = 0,
LOADING_TEXT_ENC = 1,
ENCODING_TEXT = 2,
LOADING_COND_ENC = 3,
ENCODING_COND = 4,
LOADING_DIT = 5,
LOADING_ADAPTER = 6,
ADAPTER_PRECOMPUTE = 7,
DIT_INFERENCE = 8,
LOADING_VAE = 9,
VAE_DECODE = 10,
ENCODING_OUTPUT = 11,
DONE = 12,
FAILED = 13,
CANCELLED = 14,
};
static const char * job_phase_str(JobPhase p) {
switch (p) {
case JobPhase::QUEUED: return "queued";
case JobPhase::LOADING_TEXT_ENC: return "loading_text_enc";
case JobPhase::ENCODING_TEXT: return "encoding_text";
case JobPhase::LOADING_COND_ENC: return "loading_cond_enc";
case JobPhase::ENCODING_COND: return "encoding_cond";
case JobPhase::LOADING_DIT: return "loading_dit";
case JobPhase::LOADING_ADAPTER: return "loading_adapter";
case JobPhase::ADAPTER_PRECOMPUTE: return "adapter_precompute";
case JobPhase::DIT_INFERENCE: return "dit_inference";
case JobPhase::LOADING_VAE: return "loading_vae";
case JobPhase::VAE_DECODE: return "vae_decode";
case JobPhase::ENCODING_OUTPUT: return "encoding_output";
case JobPhase::DONE: return "done";
case JobPhase::FAILED: return "failed";
case JobPhase::CANCELLED: return "cancelled";
}
return "unknown";
}
struct Job {
std::string id;
std::atomic<int> status{ 0 }; // 0=running 1=done 2=failed 3=cancelled
std::string result_body;
std::string result_mime;
std::string result_lrc; // LRC timestamp text (base64), empty if not generated
std::vector<float> result_latent; // post-DiT latent [T*64] float32, empty if not captured
std::atomic<bool> cancel{ false };
// Phase tracking (advisory, independent of `status`). phase_step/phase_total
// are optional sub-progress for phases with a natural counter; 0/0 means
// "no sub-progress available".
std::atomic<JobPhase> phase{ JobPhase::QUEUED };
std::atomic<int> phase_step{ 0 };
std::atomic<int> phase_total{ 0 };
// memory ordering contract: result_body and result_mime are written
// before status is stored (seq_cst). the client loads status (seq_cst)
// and only reads result fields after seeing done/failed. this guarantees
// visibility without an explicit mutex on the result fields.
};
// Set phase + reset sub-progress counters in one shot, at worker log-anchor
// points, so external observers never see a stale step counter from a prior
// phase.
static inline void job_set_phase(Job & job, JobPhase p, int step = 0, int total = 0) {
job.phase_step.store(step, std::memory_order_relaxed);
job.phase_total.store(total, std::memory_order_relaxed);
job.phase.store(p, std::memory_order_release);
}
static std::mutex mtx_jobs;
static std::unordered_map<std::string, std::shared_ptr<Job>> g_jobs;
static std::deque<std::string> g_job_order;
static const int MAX_JOBS = 32;
// generate a random hex ID (64 bits of entropy, non-predictable)
static std::string job_make_id() {
static std::mt19937_64 rng(std::random_device{}());
static std::mutex mtx_rng;
std::lock_guard<std::mutex> lock(mtx_rng);
char buf[17];
snprintf(buf, sizeof(buf), "%016llx", (unsigned long long) rng());
return buf;
}
static std::shared_ptr<Job> job_create() {
std::lock_guard<std::mutex> lock(mtx_jobs);
auto job = std::make_shared<Job>();
job->id = job_make_id();
g_jobs[job->id] = job;
g_job_order.push_back(job->id);
// evict oldest completed jobs to stay under MAX_JOBS.
// running jobs (status 0) are never evicted.
while ((int) g_job_order.size() > MAX_JOBS) {
bool evicted = false;
for (auto it = g_job_order.begin(); it != g_job_order.end(); ++it) {
auto jit = g_jobs.find(*it);
if (jit == g_jobs.end() || jit->second->status.load() != 0) {
if (jit != g_jobs.end()) {
g_jobs.erase(jit);
}
g_job_order.erase(it);
evicted = true;
break;
}
}
if (!evicted) {
break;
}
}
return job;
}
static std::shared_ptr<Job> job_find(const std::string & id) {
std::lock_guard<std::mutex> lock(mtx_jobs);
auto it = g_jobs.find(id);
return it != g_jobs.end() ? it->second : nullptr;
}
static const char * job_status_str(int s) {
switch (s) {
case 0:
return "running";
case 1:
return "done";
case 2:
return "failed";
case 3:
return "cancelled";
default:
return "unknown";
}
}
// log capture: intercept stderr via pipe, forward to terminal + ring buffer.
// SSE clients connect to /logs and receive lines in real time.
#define LOG_RING_BITS 9
#define LOG_RING_SIZE (1 << LOG_RING_BITS)
#define LOG_RING_MASK (LOG_RING_SIZE - 1)
static std::mutex mtx_log;
static std::condition_variable cv_log;
static std::string log_ring[LOG_RING_SIZE];
static uint64_t log_seq = 0;
static int g_real_stderr_fd = -1;
static int g_pipe_read_fd = -1;
static void setup_log_capture() {
g_real_stderr_fd = fd_dup(STDERR_FILENO);
int pipefd[2];
if (fd_pipe(pipefd) != 0) {
g_real_stderr_fd = -1;
return;
}
g_pipe_read_fd = pipefd[0];
fd_dup2(pipefd[1], STDERR_FILENO);
fd_close(pipefd[1]);
}
// reader thread: drain pipe, forward to real stderr, push lines to ring.
// exits when the write end of the pipe is closed (fd_dup2 restores real stderr).
static void log_reader_main() {
char buf[4096];
std::string partial;
for (;;) {
int n = fd_read(g_pipe_read_fd, buf, sizeof(buf));
if (n <= 0) {
break;
}
fd_write(g_real_stderr_fd, buf, (size_t) n);
partial.append(buf, (size_t) n);
size_t pos;
while ((pos = partial.find('\n')) != std::string::npos) {
std::lock_guard<std::mutex> lock(mtx_log);
log_ring[log_seq & LOG_RING_MASK] = partial.substr(0, pos);
log_seq++;
cv_log.notify_all();
partial.erase(0, pos + 1);
}
}
if (!partial.empty()) {
std::lock_guard<std::mutex> lock(mtx_log);
log_ring[log_seq & LOG_RING_MASK] = std::move(partial);
log_seq++;
cv_log.notify_all();
}
fd_close(g_pipe_read_fd);
}
static void teardown_log_capture() {
if (g_real_stderr_fd < 0) {
return;
}
fflush(stderr);
fd_dup2(g_real_stderr_fd, STDERR_FILENO);
// g_real_stderr_fd stays open: the reader thread writes to it
}
// RAII: captures stderr on construction, restores + joins reader on destruction.
// safe on any exit path (early arg errors, model load failures, normal shutdown).
struct LogCapture {
std::thread reader;
LogCapture() {
setup_log_capture();
reader = std::thread(log_reader_main);
}
~LogCapture() {
teardown_log_capture();
cv_log.notify_all();
if (reader.joinable()) {
reader.join();
}
// reader is done draining the pipe, safe to close
if (g_real_stderr_fd >= 0) {
fd_close(g_real_stderr_fd);
g_real_stderr_fd = -1;
}
}
};
// GET /logs: SSE stream of stderr lines.
// sends backlog (up to LOG_RING_SIZE) then streams new lines in real time.
static void handle_logs(const httplib::Request &, httplib::Response & res) {
res.set_header("Cache-Control", "no-cache");
res.set_header("X-Accel-Buffering", "no");
res.set_chunked_content_provider(
"text/event-stream", [cursor = uint64_t(0), init = false](size_t, httplib::DataSink & sink) mutable -> bool {
std::unique_lock<std::mutex> lock(mtx_log);
if (!init) {
uint64_t avail = log_seq < LOG_RING_SIZE ? log_seq : (uint64_t) LOG_RING_SIZE;
cursor = log_seq - avail;
while (cursor < log_seq) {
std::string ev = "data: " + log_ring[cursor & LOG_RING_MASK] + "\n\n";
cursor++;
lock.unlock();
if (!sink.write(ev.c_str(), ev.size())) {
return false;
}
lock.lock();
}
init = true;
}
cv_log.wait_for(lock, std::chrono::seconds(2));
while (cursor < log_seq) {
std::string ev = "data: " + log_ring[cursor & LOG_RING_MASK] + "\n\n";
cursor++;
lock.unlock();
if (!sink.write(ev.c_str(), ev.size())) {
return false;
}
lock.lock();
}
return true;
});
}
// cancel callback: checks the per-job cancel flag.
static bool server_cancel_job(void * data) {
auto * flag = (const std::atomic<bool> *) data;
return flag && flag->load(std::memory_order_relaxed);
}
// helper: set a JSON error response
static void json_error(httplib::Response & res, int status, const char * msg) {
yyjson_mut_doc * doc = yyjson_mut_doc_new(NULL);
yyjson_mut_val * root = yyjson_mut_obj(doc);
yyjson_mut_doc_set_root(doc, root);
yyjson_mut_obj_add_str(doc, root, "error", msg);
char * json = yyjson_mut_write(doc, 0, NULL);
yyjson_mut_doc_free(doc);
res.status = status;
res.set_content(json, "application/json");
free(json);
}
// resolve model name: explicit request > already loaded > first in bucket
static std::string resolve_name(const std::vector<ModelEntry> & bucket,
const std::string & requested,
const std::string & loaded) {
if (!requested.empty()) {
return requested;
}
if (!loaded.empty()) {
return loaded;
}
if (!bucket.empty()) {
return bucket[0].name;
}
return "";
}
// =====================================================================
// HOT-STEP EXTENSIONS
// =====================================================================
// server-side routing fields parsed from JSON (not part of AceRequest).
// these are HOT-Step additions that travel alongside the upstream request.
struct ServerFields {
std::string vae_model; // explicit VAE selection ("": use first in registry)
std::string emb_model; // explicit text encoder selection ("": use first in registry)
std::string solver_name; // "euler", "rk4", "heun", etc.
std::string scheduler; // "composite:...", "bong_tangent", etc.
std::string guidance_mode; // "apg", "dynamic_cfg", etc.
float apg_momentum = 0.75f;
float apg_norm_threshold = 2.5f;
int stork_substeps = 10;
float beat_stability = 0.25f;
float frequency_damping = 0.4f;
float temporal_smoothing = 0.13f;
AdapterGroupScales group_scales; // per-group adapter scale multipliers
std::string adapter_mode; // "merge" (default, F32 promoted) or "runtime"
std::string adapter_runtime_quant = "bf16"; // runtime delta VRAM precision: bf16/q8_0/q4_k
bool adapter_merge_lowvram = false; // merge mode: requant to native type instead of F32 promotion
float adapter_section_align_at = 0.55f; // per-section masking: alignment step fraction
float adapter_section_isolation = 0.0f; // per-section masking: cross-section self-attn penalty (0..1)
// Basin re-base: nudge adapted weights toward the adapter's training base S.
std::string rebase_source = ""; // absolute path to S (resolved by Node server)
float rebase_beta = 0.0f; // 0 = off
// DCW (Differential Correction in Wavelet domain)
bool dcw_enabled = false;
std::string dcw_mode = "low";
float dcw_scaler = 0.1f;
float dcw_high_scaler = 0.0f;
// Latent post-processing
float latent_shift = 0.0f;
float latent_rescale = 1.0f;
float cfg_cutoff_ratio = 1.0f;
float cache_ratio = 0.0f;
std::string custom_timesteps = "";
// Post-VAE spectral denoiser (HOT-Step)
float denoise_strength = 0.0f; // 0 = off, 1 = max
float denoise_smoothing = 0.7f;
float denoise_mix = 0.25f;
// Lua plugin params: {"pluginName:key": "value", ...}
std::unordered_map<std::string, std::string> plugin_params;
// Structural seed for repeated sections (Song Builder). seed_strength from
// the request JSON; seed_latents from the multipart "seed_latents" part.
float seed_strength = 0.0f;
std::vector<float> seed_latents;
// Song Builder: free the one-shot LM before this synth (repaint sections
// never use it). Only Song Builder sets this; other modes leave it false.
bool evict_lm = false;
// Per-request VRAM knobs (Song Builder / low-VRAM). 0 / -1 = loaded default.
int vae_chunk = 0; // >0: VAE tile size override
int batch_cfg = -1; // 0: split CFG, 1: batch, -1: default
};
static void parse_server_fields(const char * json, ServerFields * sf) {
sf->vae_model = "";
sf->emb_model = "";
sf->solver_name = "euler";
sf->scheduler = "";
sf->guidance_mode = "apg";
sf->adapter_mode = "merge";
sf->apg_momentum = 0.75f;
sf->apg_norm_threshold = 2.5f;
sf->stork_substeps = 10;
sf->beat_stability = 0.25f;
sf->frequency_damping = 0.4f;
sf->temporal_smoothing = 0.13f;
yyjson_doc * doc = yyjson_read(json, strlen(json), 0);
if (!doc) return;
yyjson_val * root = yyjson_doc_get_root(doc);
if (!root) { yyjson_doc_free(doc); return; }
yyjson_val * obj = root;
if (yyjson_is_arr(root)) {
obj = yyjson_arr_get_first(root);
}
if (!obj || !yyjson_is_obj(obj)) { yyjson_doc_free(doc); return; }
yyjson_val * v;
if ((v = yyjson_obj_get(obj, "vae_model")) && yyjson_is_str(v)) {
sf->vae_model = yyjson_get_str(v);
}
if ((v = yyjson_obj_get(obj, "emb_model")) && yyjson_is_str(v)) {
sf->emb_model = yyjson_get_str(v);
}
// Solver / scheduler / guidance
if ((v = yyjson_obj_get(obj, "infer_method")) && yyjson_is_str(v)) {
sf->solver_name = yyjson_get_str(v);
}
if ((v = yyjson_obj_get(obj, "scheduler")) && yyjson_is_str(v)) {
sf->scheduler = yyjson_get_str(v);
}
if ((v = yyjson_obj_get(obj, "guidance_mode")) && yyjson_is_str(v)) {
sf->guidance_mode = yyjson_get_str(v);
}
if ((v = yyjson_obj_get(obj, "adapter_mode")) && yyjson_is_str(v)) {
sf->adapter_mode = yyjson_get_str(v);
}
if ((v = yyjson_obj_get(obj, "adapter_runtime_quant")) && yyjson_is_str(v)) {
sf->adapter_runtime_quant = yyjson_get_str(v);
}
if ((v = yyjson_obj_get(obj, "adapter_merge_lowvram")) && yyjson_is_bool(v)) {
sf->adapter_merge_lowvram = yyjson_get_bool(v);
}
if ((v = yyjson_obj_get(obj, "adapter_section_align_at")) && yyjson_is_num(v)) {
sf->adapter_section_align_at = (float) yyjson_get_num(v);
}
if ((v = yyjson_obj_get(obj, "adapter_section_isolation")) && yyjson_is_num(v)) {
sf->adapter_section_isolation = (float) yyjson_get_num(v);
}
if ((v = yyjson_obj_get(obj, "rebase_source")) && yyjson_is_str(v)) {
sf->rebase_source = yyjson_get_str(v);
}
if ((v = yyjson_obj_get(obj, "rebase_beta")) && yyjson_is_num(v)) {
sf->rebase_beta = yyjson_is_real(v) ? (float) yyjson_get_real(v) : (float) yyjson_get_int(v);
}
// APG tuning
if ((v = yyjson_obj_get(obj, "apg_momentum")) && yyjson_is_num(v)) {
sf->apg_momentum = (float) yyjson_get_real(v);
}
if ((v = yyjson_obj_get(obj, "apg_norm_threshold")) && yyjson_is_num(v)) {
sf->apg_norm_threshold = (float) yyjson_get_real(v);
}
// Structural seed strength (Song Builder repeated sections)
if ((v = yyjson_obj_get(obj, "seed_strength")) && yyjson_is_num(v)) {
sf->seed_strength = (float) yyjson_get_real(v);
}
// Song Builder: evict the LM before synth (repaint sections don't use it)
if ((v = yyjson_obj_get(obj, "evict_lm")) && yyjson_is_bool(v)) {
sf->evict_lm = yyjson_get_bool(v);
}
// Per-request VRAM knobs (Song Builder / low-VRAM)
if ((v = yyjson_obj_get(obj, "vae_chunk")) && yyjson_is_int(v)) {
sf->vae_chunk = (int) yyjson_get_int(v);
}
if ((v = yyjson_obj_get(obj, "batch_cfg")) && yyjson_is_int(v)) {
sf->batch_cfg = (int) yyjson_get_int(v);
}
// STORK solver params
if ((v = yyjson_obj_get(obj, "stork_substeps")) && yyjson_is_int(v)) {
sf->stork_substeps = (int) yyjson_get_int(v);
}
if ((v = yyjson_obj_get(obj, "beat_stability")) && yyjson_is_num(v)) {
sf->beat_stability = (float) yyjson_get_real(v);
}
if ((v = yyjson_obj_get(obj, "frequency_damping")) && yyjson_is_num(v)) {
sf->frequency_damping = (float) yyjson_get_real(v);
}
if ((v = yyjson_obj_get(obj, "temporal_smoothing")) && yyjson_is_num(v)) {
sf->temporal_smoothing = (float) yyjson_get_real(v);
}
// Per-group adapter scales: {"adapter_group_scales": {"self_attn": 1.0, ...}}
// NOTE: JSON integer 1 vs float 1.0 — yyjson_get_real returns 0 for ints.
// Use a lambda that handles both.
auto get_num = [](yyjson_val * val) -> float {
return yyjson_is_real(val) ? (float) yyjson_get_real(val) : (float) yyjson_get_int(val);
};
yyjson_val * gs_obj = yyjson_obj_get(obj, "adapter_group_scales");
if (gs_obj && yyjson_is_obj(gs_obj)) {
if ((v = yyjson_obj_get(gs_obj, "self_attn")) && yyjson_is_num(v))
sf->group_scales.self_attn = get_num(v);
if ((v = yyjson_obj_get(gs_obj, "cross_attn")) && yyjson_is_num(v))
sf->group_scales.cross_attn = get_num(v);
if ((v = yyjson_obj_get(gs_obj, "mlp")) && yyjson_is_num(v))
sf->group_scales.mlp = get_num(v);
if ((v = yyjson_obj_get(gs_obj, "cond_embed")) && yyjson_is_num(v))
sf->group_scales.cond_embed = get_num(v);
if ((v = yyjson_obj_get(gs_obj, "time_embed")) && yyjson_is_num(v))
sf->group_scales.time_embed = get_num(v);
if ((v = yyjson_obj_get(gs_obj, "proj_in")) && yyjson_is_num(v))
sf->group_scales.proj_in = get_num(v);
fprintf(stderr, "[DIAG] Parsed adapter_group_scales from JSON: sa=%.2f ca=%.2f mlp=%.2f ce=%.2f te=%.2f pi=%.2f\n",
sf->group_scales.self_attn, sf->group_scales.cross_attn,
sf->group_scales.mlp, sf->group_scales.cond_embed, sf->group_scales.time_embed, sf->group_scales.proj_in);
} else {
fprintf(stderr, "[DIAG] adapter_group_scales: gs_obj=%p is_obj=%d\n",
(void*)gs_obj, gs_obj ? yyjson_is_obj(gs_obj) : -1);
}
// DCW fields
if ((v = yyjson_obj_get(obj, "dcw_enabled"))) {
if (yyjson_is_bool(v)) {
sf->dcw_enabled = yyjson_get_bool(v);
} else if (yyjson_is_str(v)) {
sf->dcw_enabled = (strcmp(yyjson_get_str(v), "true") == 0);
}
}
if ((v = yyjson_obj_get(obj, "dcw_mode")) && yyjson_is_str(v)) {
sf->dcw_mode = yyjson_get_str(v);
}
if ((v = yyjson_obj_get(obj, "dcw_scaler")) && yyjson_is_num(v)) {
sf->dcw_scaler = (float) yyjson_get_real(v);
}
if ((v = yyjson_obj_get(obj, "dcw_high_scaler")) && yyjson_is_num(v)) {
sf->dcw_high_scaler = (float) yyjson_get_real(v);
}
// Latent post-processing
if ((v = yyjson_obj_get(obj, "latent_shift")) && yyjson_is_num(v)) {
sf->latent_shift = (float) yyjson_get_real(v);
}
if ((v = yyjson_obj_get(obj, "latent_rescale")) && yyjson_is_num(v)) {
sf->latent_rescale = (float) yyjson_get_real(v);
}
if ((v = yyjson_obj_get(obj, "cfg_cutoff_ratio")) && yyjson_is_num(v)) {
sf->cfg_cutoff_ratio = (float) yyjson_get_real(v);
}
if ((v = yyjson_obj_get(obj, "cache_ratio")) && yyjson_is_num(v)) {
sf->cache_ratio = (float) yyjson_get_real(v);
}
if ((v = yyjson_obj_get(obj, "custom_timesteps")) && yyjson_is_str(v)) {
sf->custom_timesteps = yyjson_get_str(v);
}
// Post-VAE spectral denoiser (HOT-Step)
// NOTE: use get_num — JS may serialize whole numbers as integers (1 not 1.0)
if ((v = yyjson_obj_get(obj, "denoise_strength")) && yyjson_is_num(v)) {
sf->denoise_strength = get_num(v);
}
if ((v = yyjson_obj_get(obj, "denoise_smoothing")) && yyjson_is_num(v)) {
sf->denoise_smoothing = get_num(v);
}
if ((v = yyjson_obj_get(obj, "denoise_mix")) && yyjson_is_num(v)) {
sf->denoise_mix = get_num(v);
}
// Lua plugin params: iterate "plugin_params" object
yyjson_val * pp_obj = yyjson_obj_get(obj, "plugin_params");
if (pp_obj && yyjson_is_obj(pp_obj)) {
sf->plugin_params.clear();
yyjson_val * pp_key, * pp_val;
yyjson_obj_iter pp_iter;
yyjson_obj_iter_init(pp_obj, &pp_iter);
while ((pp_key = yyjson_obj_iter_next(&pp_iter))) {
pp_val = yyjson_obj_iter_get_val(pp_key);
std::string k = yyjson_get_str(pp_key);
std::string v_str;
if (yyjson_is_str(pp_val)) {
v_str = yyjson_get_str(pp_val);
} else if (yyjson_is_real(pp_val)) {
v_str = std::to_string(yyjson_get_real(pp_val));
} else if (yyjson_is_int(pp_val)) {
v_str = std::to_string(yyjson_get_int(pp_val));
} else if (yyjson_is_bool(pp_val)) {
v_str = yyjson_get_bool(pp_val) ? "true" : "false";
}
sf->plugin_params[k] = v_str;
}
if (!sf->plugin_params.empty()) {
fprintf(stderr, "[DIAG] Parsed %d plugin_params\n", (int) sf->plugin_params.size());
}
}
yyjson_doc_free(doc);
}
// =====================================================================
// Resolve a planner-LM adapter (local HOT-Step feature): registry name from
// adapters/lm/, or an explicit path (PEFT dir or .safetensors). Returns ""
// when the adapter cannot be resolved — callers must FAIL the request rather
// than silently running the base LM the user didn't ask for.
static std::string resolve_lm_adapter_path(const std::string & name_or_path) {
if (name_or_path.empty()) {
return "";
}
for (const auto & e : g_registry.lm_adapters) {
if (e.name == name_or_path) {
return e.path;
}
}
// Path fallback (absolute or relative): PEFT dir or bare safetensors file
bool looks_like_path = name_or_path.find('/') != std::string::npos ||
name_or_path.find('\\') != std::string::npos;
if (looks_like_path) {
if (registry_is_file(name_or_path.c_str())) {
return name_or_path;
}
// PEFT dir, else the LyCORIS LoKR layout ace-train --adapter-type lokr
// writes: lokr_weights.safetensors and deliberately NO
// adapter_config.json (alpha rides the per-module tensors). Probing only
// the PEFT leaf here rejected LoKR adapters before lm_adapter_load ever
// saw them, so its own fallback never ran (2026-07-30).
const char * leaves[2] = { "/adapter_model.safetensors", "/lokr_weights.safetensors" };
for (int li = 0; li < 2; li++) {
if (registry_is_file((name_or_path + leaves[li]).c_str())) {
return name_or_path;
}
}
}
fprintf(stderr, "[Server] LM adapter not found: %s (looked in adapters/lm/ registry%s)\n",
name_or_path.c_str(), looks_like_path ? " and as a path" : "");
return "";
}
// LM worker: generates metadata + lyrics + codes, stores JSON result in job.
static void lm_worker(std::shared_ptr<Job> job, AceRequest ace_req, int lm_batch_size, int mode) {
if (job->cancel.load()) {
job_set_phase(*job, JobPhase::CANCELLED);
job->status.store(3);
return;
}
// Resolve model name and build per-request params from the template.
std::string lm_name = resolve_name(g_registry.lm, ace_req.lm_model, g_loaded_lm);
const ModelEntry * entry = registry_find(g_registry.lm, lm_name.c_str());
if (!entry) {
fprintf(stderr, "[Server] LM not found: %s\n", lm_name.c_str());
job_set_phase(*job, JobPhase::FAILED);
job->status.store(2);
return;
}
AceLmParams p = g_lm_params;
p.model_path = entry->path.c_str();
// Planner-LM runtime LoRA (local HOT-Step feature)
std::string lm_adapter_path;
if (!ace_req.lm_adapter.empty()) {
lm_adapter_path = resolve_lm_adapter_path(ace_req.lm_adapter);
if (lm_adapter_path.empty()) {
job_set_phase(*job, JobPhase::FAILED);
job->status.store(2);
return;
}
p.adapter_path = lm_adapter_path.c_str();
p.adapter_scale = ace_req.lm_adapter_scale;
fprintf(stderr, "[Server] LM adapter: %s (scale %.2f)\n",
lm_adapter_path.c_str(), ace_req.lm_adapter_scale);
}
// LM has no DiT; it reuses LOADING_DIT to mean "loading the big model" so
// the wrapper has one 'loading' phase to surface across /lm and /synth.
job_set_phase(*job, JobPhase::LOADING_DIT);
// Acquire a fresh LM ctx from the shared store. Under EVICT_STRICT the
// module is reloaded if another pipeline evicted it; under EVICT_NEVER
// the store returns the cached instance.
AceLm * ctx = ace_lm_load(g_store, &p);
if (!ctx) {
fprintf(stderr, "[Server] FATAL: LM load failed\n");
job_set_phase(*job, JobPhase::FAILED);
job->status.store(2);
return;
}
job_set_phase(*job, JobPhase::DIT_INFERENCE); // LM "inference" phase
// Execute and always free the ctx, success or failure: the store decides
// whether the underlying GPU module stays resident.
// Default lm_seed to the DiT seed only when the caller didn't send an
// independent one (request_parse_json leaves it at the -1 sentinel when
// the "lm_seed" key is absent from the JSON body). This preserves the
// old single-seed behavior for clients that only send "seed" (locked
// seed -> both deterministic, random -> both random), while letting a
// client that explicitly sends "lm_seed" (e.g. a UI with independent
// LM/generation seed controls) take priority.
if (ace_req.lm_seed < 0) {
ace_req.lm_seed = ace_req.seed;
}
request_resolve_lm_seed(&ace_req);
std::vector<AceRequest> out(lm_batch_size);
int rc = ace_lm_generate(ctx, &ace_req, lm_batch_size, out.data(), NULL, NULL, server_cancel_job,
(void *) &job->cancel, mode);
ace_lm_free(ctx);
if (rc != 0) {
bool cancelled = job->cancel.load();
job_set_phase(*job, cancelled ? JobPhase::CANCELLED : JobPhase::FAILED);
job->status.store(cancelled ? 3 : 2);
return;
}
// Sticky name hint for resolve_name under --keep-loaded. Master clears it
// in the default mode since the ctx is gone; we match that behavior.
if (g_keep_loaded) {
g_loaded_lm = lm_name;
} else {
g_loaded_lm.clear();
}
// serialize output as a JSON array
std::string body = "[";
for (int i = 0; i < lm_batch_size; i++) {
if (i > 0) {
body += ",";
}
body += request_to_json(&out[i]);
}
body += "]";
job->result_body = std::move(body);
job->result_mime = "application/json";
job_set_phase(*job, JobPhase::DONE);
job->status.store(1);
fprintf(stderr, "[Server] Job %s done (LM, %d results)\n", job->id.c_str(), lm_batch_size);
}
// POST /lm
// accepts: AceRequest JSON (lm_mode in the body selects the generation mode).
// returns: JSON {"id":"N"} immediately. result is a JSON array of enriched
// AceRequests (lm_batch_size controls count).
// modes (AceRequest.lm_mode):
// generate metadata + lyrics + audio_codes (full composer pass)
// inspire metadata + lyrics (audio_codes stays empty)
// format metadata + lyrics (audio_codes stays empty)
static void handle_lm(const httplib::Request & req, httplib::Response & res) {
if (g_registry.lm.empty()) {
json_error(res, 501, "No LM models in registry");
return;
}
// Co-resident mode: flip store policy BEFORE the LM loads so it stays
// cached. Without this, gen 1 frees the LM under STRICT (the synth
// worker flips to NEVER too late), and gen 2 reloads ~8 GB on top of
// the synth models that are already resident.
const bool req_keep_loaded = req.has_param("keep_loaded") && req.get_param_value("keep_loaded") == "1";
if (req_keep_loaded && !g_keep_loaded) {
g_keep_loaded = true;
store_set_policy(g_store, EVICT_NEVER);
fprintf(stderr, "[Server] Co-resident mode activated (from /lm)\n");
}
// parse request
AceRequest ace_req;
if (!request_parse_json(&ace_req, req.body.c_str())) {
json_error(res, 400, "Invalid JSON");
return;
}
if (ace_req.caption.empty()) {
json_error(res, 400, "Caption is required");
return;
}
// Resolve lm_mode string to integer mode used by ace_lm_generate.
int mode;
if (ace_req.lm_mode == LM_MODE_NAME_GENERATE) {
mode = LM_MODE_GENERATE;
} else if (ace_req.lm_mode == LM_MODE_NAME_INSPIRE) {
mode = LM_MODE_INSPIRE;
} else if (ace_req.lm_mode == LM_MODE_NAME_FORMAT) {
mode = LM_MODE_FORMAT;
} else {
json_error(res, 400, "Invalid lm_mode (use: generate, inspire, format)");
return;
}
// clamp lm_batch_size to [1, max_batch]
int lm_batch_size = ace_req.lm_batch_size;
if (lm_batch_size < 1) {
lm_batch_size = 1;
}
if (lm_batch_size > g_max_batch) {
lm_batch_size = g_max_batch;
}
auto job = job_create();
fprintf(stderr, "[Server] Job %s created (LM, mode=%d)%s\n", job->id.c_str(), mode,
g_keep_loaded ? " [keep-loaded]" : "");
work_push([job, ace_req, lm_batch_size, mode]() { lm_worker(job, ace_req, lm_batch_size, mode); });
std::string body = "{\"id\":\"" + job->id + "\"}";
res.set_content(body, "application/json");
}
// synth worker: processes synth request, stores audio result in job.
static void synth_worker(std::shared_ptr<Job> job,
std::vector<AceRequest> ace_reqs,
ServerFields sf,
float * src_interleaved,
int src_len,
float * src_latents,
int src_T_latent,
float * ref_interleaved,
int ref_len,
float * ref_latents,
int ref_T_latent,
bool output_wav,
WavFormat wav_fmt,
int peak_clip,
bool req_keep_loaded) {
// Generate every request in one DiT batch. synth_batch_size expands each
// request into per-seed variants. Total clamped to DiT max 9.
const int batch_n = (int) ace_reqs.size();
int total_alloc = 0;
for (int ri = 0; ri < batch_n; ri++) {
int sbs = ace_reqs[ri].synth_batch_size;
total_alloc += sbs < 1 ? 1 : (sbs > 9 ? 9 : sbs);
}
if (total_alloc > 9) {
fprintf(stderr, "[Server] Batch %d exceeds DiT max 9, clamping\n", total_alloc);
total_alloc = 9;
}
std::vector<AceAudio> audio(total_alloc);
if (job->cancel.load()) {
free(src_interleaved);
free(src_latents);
free(ref_interleaved);
free(ref_latents);
job_set_phase(*job, JobPhase::CANCELLED);
job->status.store(3);
return;
}
// Resolve DiT, adapter and the text-encoder / VAE singletons.
std::string dit_name = resolve_name(g_registry.dit, ace_reqs[0].synth_model, g_loaded_dit);
const ModelEntry * dit = registry_find(g_registry.dit, dit_name.c_str());
if (!dit) {
fprintf(stderr, "[Server] DiT not found: %s\n", dit_name.c_str());
free(src_interleaved);
free(src_latents);
free(ref_interleaved);
free(ref_latents);
job->status.store(2);
return;
}
if (g_registry.text_enc.empty() || g_registry.vae.empty()) {
fprintf(stderr, "[Server] Missing Text-Enc or VAE in registry\n");
free(src_interleaved);
free(src_latents);
free(ref_interleaved);
free(ref_latents);
job->status.store(2);
return;
}
AceSynthParams p = g_synth_params;
// HOT-STEP: Text encoder model selection. Resolve by name from registry.
const ModelEntry * emb_entry = nullptr;
if (!sf.emb_model.empty()) {
emb_entry = registry_find(g_registry.text_enc, sf.emb_model.c_str());
if (!emb_entry) {
fprintf(stderr, "[Server] Text encoder not found: %s, using default\n", sf.emb_model.c_str());
}
}
p.text_encoder_path = emb_entry ? emb_entry->path.c_str() : g_registry.text_enc[0].path.c_str();
p.dit_path = dit->path.c_str();
// HOT-STEP: VAE model selection. Resolve by name from registry.
// ONNX VAE files are decoder-only — they go through the ORT decode path,
// NOT the GGML encode path. If the user selects an ONNX VAE, we route it
// to onnx_vae_path and use the first GGUF/safetensors VAE for encoding.
const ModelEntry * vae_entry = nullptr;
bool vae_is_onnx = false;
if (!sf.vae_model.empty()) {
vae_entry = registry_find(g_registry.vae, sf.vae_model.c_str());
if (!vae_entry) {
fprintf(stderr, "[Server] VAE not found: %s, using default\n", sf.vae_model.c_str());
} else if (vae_entry->name.size() >= 5 &&
vae_entry->name.substr(vae_entry->name.size() - 5) == ".onnx") {
vae_is_onnx = true;
// Route ONNX VAE to ORT decode path
p.onnx_vae_path = vae_entry->path.c_str();
fprintf(stderr, "[Server] ONNX VAE selected: %s → ORT decode path\n", vae_entry->name.c_str());
// Fall back to GGUF/safetensors for encoding
vae_entry = registry_find_non_onnx(g_registry.vae);
}
}
if (!vae_entry) {
vae_entry = registry_find_non_onnx(g_registry.vae);
}
p.vae_path = vae_entry ? vae_entry->path.c_str() : g_registry.vae[0].path.c_str();
// PP-VAE: auto-detect from registry, prefer highest precision: F32 > BF16 > F16
p.pp_vae_path = nullptr;
if (!g_registry.pp_vae.empty()) {
const char * pref[] = { "F32", "BF16", "F16" };
for (const char * tag : pref) {
for (const auto & e : g_registry.pp_vae) {
if (e.name.find(tag) != std::string::npos) {
p.pp_vae_path = e.path.c_str();
break;
}
}
if (p.pp_vae_path) break;
}
if (!p.pp_vae_path) p.pp_vae_path = g_registry.pp_vae[0].path.c_str();
}
p.adapter_path = nullptr;
p.adapter_scale = 1.0f;
// Build the adapter stack. The multi-adapter `adapters` array supersedes the
// single `adapter` field; when only the single field is set we fold it into a
// one-element stack so merge/runtime loading takes a single code path. The
// resolved paths live in g_hotstep_params.adapters (read by dit_ggml_load);
// p.adapter_path points at the primary so the single-adapter gate stays armed.
g_hotstep_params.adapters.clear();
{
std::vector<AceAdapterRef> req_adapters = ace_reqs[0].adapters;
if (req_adapters.empty() && !ace_reqs[0].adapter.empty()) {
req_adapters.push_back({ ace_reqs[0].adapter, ace_reqs[0].adapter_scale });
}
for (const auto & ar : req_adapters) {
std::string path;
const AdapterEntry * adapter = registry_find_adapter(g_registry, ar.name.c_str());
if (adapter) {
path = adapter->path;
} else {
// HOT-STEP: absolute-path fallback for adapters not in the registry
// — a bare .safetensors file, or a PEFT DIRECTORY (the per-base
// adapter layout stores every trained adapter as one). fopen()
// fails on a directory, so probe the canonical weights file
// inside it and pass the DIRECTORY through — that is the shape
// registry_scan_adapters() produces, which the merge/runtime
// loaders are proven on. Mirrors resolve_lm_adapter_path().
FILE * test = fopen(ar.name.c_str(), "rb");
if (test) {
fclose(test);
fprintf(stderr, "[Server] Adapter absolute path: %s\n", ar.name.c_str());
path = ar.name;
} else {
// PEFT first, then the LyCORIS LoKR layout that ace-train
// --adapter-type lokr writes (lokr_weights.safetensors, no
// adapter_model.safetensors).
const char * leaves[2] = { "/adapter_model.safetensors", "/lokr_weights.safetensors" };
for (int li = 0; li < 2 && path.empty(); li++) {
const std::string inner_path = ar.name + leaves[li];
FILE * inner = fopen(inner_path.c_str(), "rb");
if (inner) {
fclose(inner);
fprintf(stderr, "[Server] Adapter dir (%s): %s\n", leaves[li] + 1, ar.name.c_str());
path = ar.name;
}
}
}
}
if (path.empty()) {
fprintf(stderr, "[Server] Adapter not found: %s\n", ar.name.c_str());
free(src_interleaved);
free(src_latents);
free(ref_interleaved);
free(ref_latents);
job->status.store(2);
return;
}
g_hotstep_params.adapters.push_back({ path, ar.scale, ar.gain_curve, ar.gain_in_steps });
}
if (!g_hotstep_params.adapters.empty()) {
p.adapter_path = g_hotstep_params.adapters[0].path.c_str();
p.adapter_scale = g_hotstep_params.adapters[0].scale;
}
}
fprintf(stderr, "[Server] Text encoder: %s\n", emb_entry ? sf.emb_model.c_str() : g_registry.text_enc[0].name.c_str());
fprintf(stderr, "[Server] Loading synth: DiT=%s%s%s\n", dit_name.c_str(),
g_hotstep_params.adapters.empty() ? "" : " Adapters=",
(g_keep_loaded || req_keep_loaded) ? " [keep-loaded]" : "");
for (const auto & a : g_hotstep_params.adapters) {
if (a.gain_curve.empty()) {
fprintf(stderr, "[Server] adapter: %s (scale=%.2f)\n", a.path.c_str(), a.scale);
} else {
fprintf(stderr, "[Server] adapter: %s (scale=%.2f, gain curve %zu pts, g(1)=%.2f g(0.5)=%.2f g(0)=%.2f)\n",
a.path.c_str(), a.scale, a.gain_curve.size(),
hotstep_adapter_gain(a.gain_curve, 1.0f),
hotstep_adapter_gain(a.gain_curve, 0.5f),
hotstep_adapter_gain(a.gain_curve, 0.0f));
}
}
// HOT-STEP: per-request co-resident mode. Once flipped to NEVER, stays
// that way until restart (going back to STRICT would need a full eviction
// pass and is not safe mid-flight).
if (req_keep_loaded && !g_keep_loaded) {
g_keep_loaded = true;
store_set_policy(g_store, EVICT_NEVER);
}
// HOT-Step: Song Builder frees the one-shot LM before loading the synth
// pipeline — its repaint sections never use the LM, so under keep-loaded it
// would otherwise sit in VRAM all session. Targeted (LM only); only Song
// Builder sets evict_lm, so other generation modes are unaffected.
if (sf.evict_lm) {
store_evict_lm(g_store);
}
// HOT-Step sideband: push custom params to global BEFORE synth load.
// Critical: adapter_group_scales must be set before ace_synth_load()
// because the adapter merge (inside dit_ggml_load) reads them from the
// global at merge time. Setting them after load uses stale scales.
g_hotstep_params.solver_name = sf.solver_name;
g_hotstep_params.scheduler = sf.scheduler;
g_hotstep_params.guidance_mode = sf.guidance_mode;
g_hotstep_params.apg_momentum = sf.apg_momentum;
g_hotstep_params.apg_norm_threshold = sf.apg_norm_threshold;
g_hotstep_params.stork_substeps = sf.stork_substeps;
g_hotstep_params.beat_stability = sf.beat_stability;
g_hotstep_params.frequency_damping = sf.frequency_damping;
g_hotstep_params.temporal_smoothing = sf.temporal_smoothing;
g_hotstep_params.adapter_group_scales = sf.group_scales;
g_hotstep_params.adapter_mode = sf.adapter_mode;
g_hotstep_params.adapter_runtime_quant = sf.adapter_runtime_quant.empty() ? "bf16" : sf.adapter_runtime_quant;
g_hotstep_params.adapter_merge_lowvram = sf.adapter_merge_lowvram;
g_hotstep_params.adapter_section_align_at = sf.adapter_section_align_at;
g_hotstep_params.adapter_section_isolation = sf.adapter_section_isolation;
// Per-section adapter masking (regional LoRA). Carry the parsed sections into
// the sideband and force runtime mode — merge bakes weights and cannot vary
// per frame. Active with a multi-adapter stack, or with any stack (even a
// single adapter) whose entries carry timestep gain curves — step gating
// rides the same per-adapter mask machinery via a synthetic single section.
g_hotstep_params.adapter_sections.clear();
if (!ace_reqs[0].adapter_sections.empty() &&
(g_hotstep_params.adapters.size() >= 2 ||
(!g_hotstep_params.adapters.empty() && hotstep_adapter_gains_active(g_hotstep_params.adapters)))) {
for (const auto & s : ace_reqs[0].adapter_sections) {
AdapterSection sec;
sec.weights = s.weights;
sec.size = s.size;
g_hotstep_params.adapter_sections.push_back(sec);
}
if (g_hotstep_params.adapter_mode != "runtime") {
fprintf(stderr, "[Adapter] Per-section masking active — forcing runtime mode\n");
g_hotstep_params.adapter_mode = "runtime";
}
fprintf(stderr, "[Adapter] Per-section masking: %zu sections over %zu adapters\n",
g_hotstep_params.adapter_sections.size(), g_hotstep_params.adapters.size());
}
// Basin re-base: sf.rebase_source is a DiT model NAME (same ids as the model
// selector); resolve it to its on-disk path. Must be a safetensors model dir
// (or model.safetensors) for the nudge to read F32 weights — GGUF-only sources
// fail st_open in adapter_merge and the nudge is skipped with a warning.
g_hotstep_params.rebase_beta = sf.rebase_beta;
g_hotstep_params.rebase_source = "";
if (!sf.rebase_source.empty() && sf.rebase_beta != 0.0f) {
const ModelEntry * rb = registry_find(g_registry.dit, sf.rebase_source.c_str());
if (rb) {
g_hotstep_params.rebase_source = rb->path;
fprintf(stderr, "[Server] Basin re-base: source=%s (%s), beta=%.2f\n",
sf.rebase_source.c_str(), rb->path.c_str(), sf.rebase_beta);
} else {
fprintf(stderr, "[Server] WARNING: basin re-base source model not found: %s (skipping nudge)\n",
sf.rebase_source.c_str());
}
}
g_hotstep_params.dcw_enabled = sf.dcw_enabled;
g_hotstep_params.dcw_mode = sf.dcw_mode;
g_hotstep_params.dcw_scaler = sf.dcw_scaler;
g_hotstep_params.dcw_high_scaler = sf.dcw_high_scaler;
g_hotstep_params.latent_shift = sf.latent_shift;
g_hotstep_params.latent_rescale = sf.latent_rescale;
g_hotstep_params.custom_timesteps = sf.custom_timesteps;
g_hotstep_params.cfg_cutoff_ratio = sf.cfg_cutoff_ratio;
g_hotstep_params.cache_ratio = sf.cache_ratio;
g_hotstep_params.plugin_params = sf.plugin_params;
g_hotstep_params.seed_strength = sf.seed_strength;
g_hotstep_params.seed_latents = sf.seed_latents;
g_hotstep_params.vae_chunk_override = sf.vae_chunk;
g_hotstep_params.batch_cfg_override = sf.batch_cfg;
fprintf(stderr, "[Server] HOT-Step params: solver=%s, guidance=%s, scheduler=%s\n",
sf.solver_name.c_str(), sf.guidance_mode.c_str(),
sf.scheduler.empty() ? "(default)" : sf.scheduler.c_str());
fprintf(stderr, "[Server] Adapter group scales: self_attn=%.2f, cross_attn=%.2f, mlp=%.2f, cond_embed=%.2f\n",
sf.group_scales.self_attn, sf.group_scales.cross_attn,
sf.group_scales.mlp, sf.group_scales.cond_embed);
if (sf.dcw_enabled) {
fprintf(stderr, "[Server] DCW: mode=%s scaler=%.3f high_scaler=%.3f\n",
sf.dcw_mode.c_str(), sf.dcw_scaler, sf.dcw_high_scaler);
}
if (sf.cfg_cutoff_ratio < 1.0f) {
fprintf(stderr, "[Server] CFG cutoff: ratio=%.2f (CFG for first %.0f%% of steps)\n",
sf.cfg_cutoff_ratio, sf.cfg_cutoff_ratio * 100.0f);
}
if (sf.cache_ratio > 0.0f) {
fprintf(stderr, "[Server] Step cache: ratio=%.2f (skip ~%.0f%% of forward passes)\n",
sf.cache_ratio, sf.cache_ratio * 100.0f);
}
// ace_synth_load fans out into text-enc + cond-enc + DiT + adapter (LoKr
// precompute) + VAE setup, all serialized. Without per-sub-load callbacks we
// mark the whole call as the heaviest phase: ADAPTER_PRECOMPUTE when an
// adapter is in play (the ~17 s stall), else LOADING_DIT. The wrapper keys
// on this to explain why a job is silent for 15+ s with no DiT step logs.
job_set_phase(*job, g_hotstep_params.adapters.empty() ? JobPhase::LOADING_DIT
: JobPhase::ADAPTER_PRECOMPUTE);
// Wire the per-job cancel flag into the adapter precompute loops so a cancel
// during cold start aborts in <100 ms instead of waiting for all deltas.
// Cleared via RAII on every exit path below.
g_adapter_cancel.store(&job->cancel, std::memory_order_release);
struct AdapterCancelGuard {
~AdapterCancelGuard() { g_adapter_cancel.store(nullptr, std::memory_order_release); }
} adapter_cancel_guard;
AceSynth * ctx = ace_synth_load(g_store, &p);
if (!ctx) {
fprintf(stderr, "[Server] FATAL: synth load failed\n");
free(src_interleaved);
free(src_latents);
free(ref_interleaved);
free(ref_latents);
bool cancelled = job->cancel.load();
job_set_phase(*job, cancelled ? JobPhase::CANCELLED : JobPhase::FAILED);
job->status.store(cancelled ? 3 : 2);
return;
}
job_set_phase(*job, JobPhase::DIT_INFERENCE, 0, ace_reqs[0].inference_steps);
// HOT-Step: restore auto-shift that upstream removed.
// When shift == -1, compute adaptive shift from duration + step count.
// base_shift=3.0 always — merge/turbo models need high shift.
// Upstream treats shift <= 0 as "default" (1.0 for non-turbo), which is wrong for our models.
// Must run BEFORE groups are built (copies are taken below).
for (int ri = 0; ri < batch_n; ri++) {
if (ace_reqs[ri].shift == -1.0f) {
float dur = ace_reqs[ri].duration > 0.0f ? (float) ace_reqs[ri].duration : 60.0f;
int steps = ace_reqs[ri].inference_steps > 0 ? ace_reqs[ri].inference_steps : 20;
float dur_f = 1.0f + 0.15f * ((dur - 60.0f) / 60.0f);
dur_f = fmaxf(0.8f, fminf(1.4f, dur_f));
float step_f = 1.0f + 0.1f * ((30.0f - (float) steps) / 30.0f);
step_f = fmaxf(0.8f, fminf(1.4f, step_f));
float computed = fmaxf(1.0f, fminf(6.0f, 3.0f * dur_f * step_f));
ace_reqs[ri].shift = computed;
if (ri == 0) {
fprintf(stderr, "[Server] Auto shift: duration=%.0fs, steps=%d → shift=%.3f\n",
dur, steps, computed);
}
}
}
// Build the flat batch. Seeds are resolved per original request, then
// synth_batch_size is expanded into per-seed variants in groups[0].
std::vector<std::vector<AceRequest>> groups(1);
groups[0].reserve(total_alloc);
int off = 0;
for (int ri = 0; ri < batch_n && off < total_alloc; ri++) {
auto & r = ace_reqs[ri];
int sbs = r.synth_batch_size;
if (sbs < 1) {
sbs = 1;
}
if (sbs > 9) {
sbs = 9;
}
if (off + sbs > total_alloc) {
sbs = total_alloc - off;
}
request_resolve_seed(&r);
const long long base_seed = r.seed;
for (int i = 0; i < sbs; i++) {
AceRequest v = r;
v.seed = base_seed + i;
groups[0].push_back(v);
}
off += sbs;
}
if (total_alloc > 1) {
fprintf(stderr, "[Server] Batch: %d track(s) from %d request(s)\n", total_alloc, batch_n);
}
// Two-phase run (+ optional Phase 3 LRC).
std::vector<std::string> lrc_results(total_alloc);
std::vector<std::vector<float>> captured_latents;
const int rc = synth_batch_run(ctx, groups,
src_interleaved, src_len,
src_latents, src_T_latent,
ref_interleaved, ref_len,
ref_latents, ref_T_latent,
audio.data(),
lrc_results.data(),
&captured_latents,
server_cancel_job, (void *) &job->cancel);
ace_synth_free(ctx);
free(src_interleaved);
free(src_latents);
free(ref_interleaved);
free(ref_latents);
// Store first track's post-DiT latent for retrieval via /job?latent=1
if (!captured_latents.empty() && !captured_latents[0].empty()) {
job->result_latent = std::move(captured_latents[0]);
fprintf(stderr, "[Server] Latent captured: T=%zu (%.1fs @ 25Hz)\n",
job->result_latent.size() / 64, (float)(job->result_latent.size() / 64) / 25.0f);
}
// Store LRC for the first track (used by the Node server)
if (!lrc_results.empty() && !lrc_results[0].empty()) {
// Base64 encode the LRC text for safe transport in HTTP header
const std::string & lrc = lrc_results[0];
static const char b64[] = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
std::string encoded_lrc;
encoded_lrc.reserve((lrc.size() + 2) / 3 * 4);
for (size_t i = 0; i < lrc.size(); i += 3) {
uint32_t v = ((uint8_t)lrc[i]) << 16;
if (i + 1 < lrc.size()) v |= ((uint8_t)lrc[i + 1]) << 8;
if (i + 2 < lrc.size()) v |= ((uint8_t)lrc[i + 2]);
encoded_lrc += b64[(v >> 18) & 0x3F];
encoded_lrc += b64[(v >> 12) & 0x3F];
encoded_lrc += (i + 1 < lrc.size()) ? b64[(v >> 6) & 0x3F] : '=';
encoded_lrc += (i + 2 < lrc.size()) ? b64[v & 0x3F] : '=';
}
job->result_lrc = encoded_lrc;
fprintf(stderr, "[Server] LRC: %zu bytes raw, %zu base64\n", lrc.size(), encoded_lrc.size());
}
if (rc != 0) {
for (auto & a : audio) {
ace_audio_free(&a);
}
bool cancelled = job->cancel.load();
job_set_phase(*job, cancelled ? JobPhase::CANCELLED : JobPhase::FAILED);
job->status.store(cancelled ? 3 : 2);
return;
}
// Sticky name hints for resolve_name under --keep-loaded. Master clears
// them in the default mode since the ctx is gone; we match that behavior.
if (g_keep_loaded) {
g_loaded_dit = dit_name;
g_loaded_adapter = ace_reqs[0].adapter;
g_loaded_adapter_scale = ace_reqs[0].adapter_scale;
g_loaded_vae = sf.vae_model;
} else {
g_loaded_dit.clear();
g_loaded_adapter.clear();
g_loaded_adapter_scale = 1.0f;
g_loaded_vae.clear();
}
const int total_tracks = total_alloc;
// VAE decode happened inside synth_batch_run; now peak-normalize + encode.
job_set_phase(*job, JobPhase::ENCODING_OUTPUT, 0, total_tracks);
// encode each track (peak normalize + encode)
const char * mime = output_wav ? "audio/wav" : "audio/mpeg";
std::vector<std::string> encoded(total_tracks);
for (int b = 0; b < total_tracks; b++) {
if (!audio[b].samples) {
continue;
}
// Normalize first: the noise profile was computed from normalized audio
// (peak ≈ 1.0), so the denoiser must run at normalized levels to match.
if (!output_wav || wav_fmt != WAV_F32) {
audio_normalize(audio[b].samples, audio[b].n_samples * 2, peak_clip);
}
// HOT-Step: Post-VAE spectral denoiser. Runs on the normalized planar
// stereo buffer to remove VAE fuzz/fizz using the noise profile.
if (sf.denoise_strength > 0.0f) {
audio_denoise(audio[b].samples, audio[b].n_samples, 48000,
sf.denoise_strength, sf.denoise_smoothing, sf.denoise_mix,
g_noise_profile.valid ? &g_noise_profile : nullptr);
}
if (output_wav) {
encoded[b] = audio_encode_wav(audio[b].samples, audio[b].n_samples, 48000, wav_fmt);
} else {
encoded[b] = audio_encode_mp3(audio[b].samples, audio[b].n_samples, 48000, g_mp3_kbps, server_cancel_job,
(void *) &job->cancel);
}
ace_audio_free(&audio[b]);
}
// store result in job
// single track: raw audio body
if (total_tracks == 1) {
job->result_body = std::move(encoded[0]);
job->result_mime = mime;
} else {
// multiple tracks: multipart/mixed, each part is raw audio
std::string boundary = "ace-batch-boundary";
std::string body;
for (int b = 0; b < total_tracks; b++) {
body += "--" + boundary + "\r\n";
body += "Content-Type: ";
body += mime;
body += "\r\n\r\n";
body += encoded[b];
body += "\r\n";
}
body += "--" + boundary + "--\r\n";
job->result_body = std::move(body);
job->result_mime = "multipart/mixed; boundary=" + boundary;
}
bool cancelled = job->cancel.load();
job_set_phase(*job, cancelled ? JobPhase::CANCELLED : JobPhase::DONE);
job->status.store(cancelled ? 3 : 1);
fprintf(stderr, "[Server] Job %s done (%d tracks)\n", job->id.c_str(), total_tracks);
}
// POST /synth[?format=wav16|wav24|wav32]
// returns JSON {"id":"N"} immediately.
// input:
// application/json body -> single request {} or batch [{req0}, {req1}, ...]
// multipart/form-data -> single request + audio file(s)
// part "request": JSON text
// part "audio": source audio (WAV or MP3)
// part "ref_audio": timbre reference audio (WAV or MP3), optional
// output: audio/mpeg (default) or audio/wav (?format=wav16|wav24|wav32)
// batch == 1: raw audio body
// batch > 1: multipart/mixed, each part is raw audio
// Batch size = number of JSON objects (after synth_batch_size expansion, clamped to 9).
// Metadata (seed, duration, etc) is already in the request JSON from /lm.
static void handle_synth(const httplib::Request & req, httplib::Response & res) {
if (g_registry.dit.empty() || g_registry.text_enc.empty() || g_registry.vae.empty()) {
json_error(res, 501, "No synth models in registry (need dit + text-encoder + vae)");
return;
}
// parse HOT-Step server fields (vae_model) from JSON body
ServerFields sf;
// parse request: plain JSON (single or array) or multipart (JSON + audio file).
// synth_model, lm_model, adapter, adapter_scale travel inside AceRequest now.
std::vector<AceRequest> ace_reqs;
float * src_interleaved = nullptr;
int src_len = 0;
float * src_latents = nullptr;
int src_T_latent = 0;
float * ref_interleaved = nullptr;
int ref_len = 0;
float * ref_latents = nullptr;
int ref_T_latent = 0;
if (req.is_multipart_form_data()) {
// multipart mode: single request + optional audio files
AceRequest ace_req;
std::string json_body;
if (req.form.has_file("request")) {
json_body = req.form.get_file("request").content;
} else if (req.form.has_field("request")) {
json_body = req.form.get_field("request");
} else {
json_error(res, 400, "Multipart: missing 'request' part");
return;
}
parse_server_fields(json_body.c_str(), &sf);
if (!request_parse_json(&ace_req, json_body.c_str())) {
json_error(res, 400, "Multipart: invalid JSON in 'request' part");
return;
}
if (req.form.has_file("audio")) {
auto file = req.form.get_file("audio");
if (file.content.empty()) {
json_error(res, 400, "Multipart: empty 'audio' part");
return;
}
int T_audio = 0;
float * planar = audio_read_48k_buf((const uint8_t *) file.content.data(), file.content.size(), &T_audio);
if (!planar || T_audio <= 0) {
json_error(res, 400, "Failed to decode audio");
return;
}
fprintf(stderr, "[Server] Source audio: %.2fs @ 48kHz\n", (float) T_audio / 48000.0f);
src_interleaved = audio_planar_to_interleaved(planar, T_audio);
free(planar);
src_len = T_audio;
}
if (req.form.has_file("ref_audio")) {
auto file = req.form.get_file("ref_audio");
if (!file.content.empty()) {
int T_audio = 0;
float * planar =
audio_read_48k_buf((const uint8_t *) file.content.data(), file.content.size(), &T_audio);
if (planar && T_audio > 0) {
fprintf(stderr, "[Server] Reference audio: %.2fs @ 48kHz\n", (float) T_audio / 48000.0f);
ref_interleaved = audio_planar_to_interleaved(planar, T_audio);
free(planar);
ref_len = T_audio;
} else {
fprintf(stderr, "[Server] WARNING: failed to decode ref_audio, ignoring\n");
}
}
}
// Source latents (raw float32, alternative to source audio — skips VAE encode)
if (req.form.has_file("src_latents")) {
auto file = req.form.get_file("src_latents");
if (!file.content.empty()) {
if (file.content.size() % (64 * sizeof(float)) != 0) {
json_error(res, 400, "src_latents size must be a multiple of 256 bytes (64 * float32)");
return;
}
src_T_latent = (int)(file.content.size() / (64 * sizeof(float)));
src_latents = (float *) malloc(file.content.size());
memcpy(src_latents, file.content.data(), file.content.size());
fprintf(stderr, "[Server] Source latents: T=%d (%.2fs @ 25Hz)\n",
src_T_latent, (float)src_T_latent / 25.0f);
}
}
// Reference latents (raw float32, alternative to ref audio — skips timbre VAE encode)
if (req.form.has_file("ref_latents")) {
auto file = req.form.get_file("ref_latents");
if (!file.content.empty()) {
if (file.content.size() % (64 * sizeof(float)) != 0) {
json_error(res, 400, "ref_latents size must be a multiple of 256 bytes (64 * float32)");
return;
}
ref_T_latent = (int)(file.content.size() / (64 * sizeof(float)));
ref_latents = (float *) malloc(file.content.size());
memcpy(ref_latents, file.content.data(), file.content.size());
fprintf(stderr, "[Server] Reference latents: T=%d (%.2fs @ 25Hz)\n",
ref_T_latent, (float)ref_T_latent / 25.0f);
}
}
// Structural seed latents (raw float32) — bias the repaint region's init
// noise toward an earlier section (Song Builder repeated sections).
if (req.form.has_file("seed_latents")) {
auto file = req.form.get_file("seed_latents");
if (!file.content.empty()) {
if (file.content.size() % (64 * sizeof(float)) != 0) {
json_error(res, 400, "seed_latents size must be a multiple of 256 bytes (64 * float32)");
return;
}
int seed_T = (int)(file.content.size() / (64 * sizeof(float)));
const float * sp = reinterpret_cast<const float *>(file.content.data());
sf.seed_latents.assign(sp, sp + (size_t) seed_T * 64);
fprintf(stderr, "[Server] Seed latents: T=%d (%.2fs @ 25Hz)\n", seed_T, (float)seed_T / 25.0f);
}
}
ace_reqs.push_back(ace_req);
} else {
// plain JSON body: single object {} or array [{}, ...]
fprintf(stderr, "[DIAG] /synth body (first 300 chars): %.300s\n", req.body.c_str());
parse_server_fields(req.body.c_str(), &sf);
if (!request_parse_json_array(req.body.c_str(), &ace_reqs)) {
json_error(res, 400, "Invalid JSON");
return;
}
}
if (ace_reqs.empty()) {
json_error(res, 400, "Empty request");
return;
}
if (ace_reqs[0].caption.empty() && ace_reqs[0].task_type != TASK_LEGO && ace_reqs[0].task_type != TASK_EXTRACT &&
ace_reqs[0].task_type != TASK_COMPLETE && ace_reqs[0].task_type != TASK_COVER && ace_reqs[0].task_type != TASK_REPAINT) {
json_error(res, 400, "Caption is required");
return;
}
// HOT-STEP: Output format from URL ?format= param (backward compat with our Node.js)
// Falls back to AceRequest.output_format if URL param not present.
bool output_wav = false;
WavFormat wav_fmt = WAV_S16;
{
std::string fmt_str;
if (req.has_param("format")) {
fmt_str = req.get_param_value("format");
} else {
fmt_str = ace_reqs[0].output_format;
}
bool is_mp3 = true;
if (!audio_parse_format(fmt_str.c_str(), is_mp3, wav_fmt)) {
json_error(res, 400, "Invalid format (use: mp3, wav16, wav24, wav32)");
return;
}
output_wav = !is_mp3;
}
int peak_clip = ace_reqs[0].peak_clip;
// create job, spawn worker, return ID
auto job = job_create();
fprintf(stderr, "[Server] Job %s created (%d requests)\n", job->id.c_str(), (int) ace_reqs.size());
// per-request co-resident mode: ?keep_loaded=1
const bool req_keep_loaded = req.has_param("keep_loaded") && req.get_param_value("keep_loaded") == "1";
work_push([job, reqs = std::move(ace_reqs), sf, src_interleaved, src_len, src_latents, src_T_latent,
ref_interleaved, ref_len, ref_latents, ref_T_latent, output_wav, wav_fmt, peak_clip, req_keep_loaded]() mutable {
synth_worker(job, std::move(reqs), sf, src_interleaved, src_len, src_latents, src_T_latent,
ref_interleaved, ref_len, ref_latents, ref_T_latent, output_wav, wav_fmt, peak_clip, req_keep_loaded);
});
// return job ID immediately
std::string body = "{\"id\":\"" + job->id + "\"}";
res.set_content(body, "application/json");
}
// understand worker: load LM + tokenizer, run understand, store JSON result in job.
static void understand_worker(std::shared_ptr<Job> job, AceRequest ace_req, float * src_interleaved, int src_len) {
if (job->cancel.load()) {
free(src_interleaved);
job->status.store(3);
return;
}
// Resolve LM + DiT (the DiT path carries the tokenizer weights).
std::string lm_name = resolve_name(g_registry.lm, ace_req.lm_model, g_loaded_lm);
std::string dit_name = resolve_name(g_registry.dit, ace_req.synth_model, g_loaded_dit);
const ModelEntry * lm_entry = registry_find(g_registry.lm, lm_name.c_str());
const ModelEntry * dit = registry_find(g_registry.dit, dit_name.c_str());
if (!lm_entry || !dit) {
fprintf(stderr, "[Server] LM or DiT not found: lm=%s dit=%s\n", lm_name.c_str(), dit_name.c_str());
free(src_interleaved);
job->status.store(2);
return;
}
AceUnderstandParams p = g_und_params;
p.model_path = lm_entry->path.c_str();
p.dit_path = dit->path.c_str();
AceUnderstand * ctx = ace_understand_load(g_store, &p);
if (!ctx) {
fprintf(stderr, "[Server] FATAL: understand load failed\n");
free(src_interleaved);
job->status.store(2);
return;
}
AceRequest out;
int rc = ace_understand_generate(ctx, src_interleaved, src_len,
nullptr, 0, // src_latents (audio path)
&ace_req, &out,
nullptr, nullptr, // latent_out, T_latent_out
server_cancel_job, (void *) &job->cancel);
ace_understand_free(ctx);
free(src_interleaved);
if (rc != 0) {
job->status.store(job->cancel.load() ? 3 : 2);
return;
}
// Sticky name hints for resolve_name under --keep-loaded. Master clears
// them in the default mode since the ctx is gone; we match that behavior.
if (g_keep_loaded) {
g_loaded_lm = lm_name;
g_loaded_und_dit = dit_name;
} else {
g_loaded_lm.clear();
g_loaded_und_dit.clear();
}
job->result_body = "[" + request_to_json(&out) + "]";
job->result_mime = "application/json";
job->status.store(1);
fprintf(stderr, "[Server] Job %s done (understand)\n", job->id.c_str());
}
// POST /understand
// multipart/form-data: full pipeline (audio + optional JSON params)
// part "audio": WAV or MP3 file (required)
// part "request": JSON text (optional, for model selection and sampling params)
// returns: JSON {"id":"N"} immediately.
static void handle_understand(const httplib::Request & req, httplib::Response & res) {
if (g_registry.lm.empty() || g_registry.dit.empty() || g_registry.vae.empty()) {
json_error(res, 501, "Understand requires LM, DiT and VAE models");
return;
}
if (!req.is_multipart_form_data()) {
json_error(res, 400, "Understand requires multipart/form-data");
return;
}
// parse multipart: required "audio" part, optional "request" part for sampling params.
// synth_model, lm_model, adapter, adapter_scale travel inside AceRequest.
AceRequest ace_req;
request_init(&ace_req);
ace_req.lm_temperature = 0.3f; // understand default: lower than generation
ace_req.lm_top_p = 1.0f; // understand default: no nucleus sampling
if (req.form.has_file("request")) {
const std::string & json = req.form.get_file("request").content;
if (!request_parse_json(&ace_req, json.c_str())) {
json_error(res, 400, "Multipart: invalid JSON in 'request' part");
return;
}
} else if (req.form.has_field("request")) {
const std::string & json = req.form.get_field("request");
if (!request_parse_json(&ace_req, json.c_str())) {
json_error(res, 400, "Multipart: invalid JSON in 'request' part");
return;
}
}
if (!req.form.has_file("audio")) {
json_error(res, 400, "Multipart: missing 'audio' part");
return;
}
auto file = req.form.get_file("audio");
if (file.content.empty()) {
json_error(res, 400, "Multipart: empty 'audio' part");
return;
}
// decode directly from multipart buffer (WAV/MP3 auto-detected)
int T_audio = 0;
float * planar = audio_read_48k_buf((const uint8_t *) file.content.data(), file.content.size(), &T_audio);
if (!planar || T_audio <= 0) {
json_error(res, 400, "Failed to decode audio");
return;
}
fprintf(stderr, "[Server] Understand source: %.2fs @ 48kHz\n", (float) T_audio / 48000.0f);
// convert planar [L:T][R:T] to interleaved [L0,R0,L1,R1,...] for pipeline
float * src_interleaved = audio_planar_to_interleaved(planar, T_audio);
free(planar);
int src_len = T_audio;
auto job = job_create();
fprintf(stderr, "[Server] Job %s created (understand)\n", job->id.c_str());
work_push(
[job, ace_req, src_interleaved, src_len]() { understand_worker(job, ace_req, src_interleaved, src_len); });
std::string body = "{\"id\":\"" + job->id + "\"}";
res.set_content(body, "application/json");
}
// ────────────────────────────────────────────────────────────────────────
// /vae endpoint: standalone VAE encode/decode (ported from upstream)
// ────────────────────────────────────────────────────────────────────────
// decode worker: VAE decode only. Loads the requested VAE decoder,
// decodes latents to 48kHz stereo audio, encodes to requested format,
// stores in job. Client already holds the latents it sent.
static void vae_decode_worker(std::shared_ptr<Job> job,
AceRequest ace_req,
std::vector<float> src_latents,
int src_T_latent,
bool output_wav,
WavFormat wav_fmt,
int peak_clip) {
if (job->cancel.load()) {
job->status.store(3);
return;
}
std::string vae_name = resolve_name(g_registry.vae, ace_req.vae, g_loaded_vae);
const ModelEntry * vae_entry = registry_find(g_registry.vae, vae_name.c_str());
if (!vae_entry) {
fprintf(stderr, "[Server] decode: VAE not found: %s\n", vae_name.c_str());
job->status.store(2);
return;
}
ModelKey vae_key;
vae_key.kind = MODEL_VAE_DEC;
vae_key.path = vae_entry->path;
vae_key.adapter_scale = 1.0f;
auto t_start = std::chrono::steady_clock::now();
VAEGGML * vae = store_require_vae_dec(g_store, vae_key);
if (!vae) {
fprintf(stderr, "[Server] decode: store_require_vae_dec failed\n");
job->status.store(2);
return;
}
ModelHandle vae_guard(g_store, vae);
int T_audio_max = (src_T_latent + 64) * 1920;
std::vector<float> audio_buf((size_t) T_audio_max * 2);
int T_audio = vae_ggml_decode_tiled(vae, src_latents.data(), src_T_latent, audio_buf.data(), T_audio_max,
g_synth_params.vae_chunk, g_synth_params.vae_overlap);
if (T_audio < 0) {
fprintf(stderr, "[Server] decode: vae_ggml_decode_tiled failed\n");
job->status.store(2);
return;
}
auto t_end = std::chrono::steady_clock::now();
float ms = (float) std::chrono::duration_cast<std::chrono::microseconds>(t_end - t_start).count() / 1000.0f;
fprintf(stderr, "[Server] decode: %d latent frames -> %d audio samples (%.2fs), %.0fms\n", src_T_latent, T_audio,
(float) T_audio / 48000.0f, ms);
if (g_keep_loaded) {
g_loaded_vae = vae_name;
} else {
g_loaded_vae.clear();
}
if (!output_wav || wav_fmt != WAV_F32) {
audio_normalize(audio_buf.data(), T_audio * 2, peak_clip);
}
std::string encoded;
const char * mime = output_wav ? "audio/wav" : "audio/mpeg";
if (output_wav) {
encoded = audio_encode_wav(audio_buf.data(), T_audio, 48000, wav_fmt);
} else {
encoded = audio_encode_mp3(audio_buf.data(), T_audio, 48000, ace_req.mp3_bitrate, server_cancel_job,
(void *) &job->cancel);
}
job->result_body = std::move(encoded);
job->result_mime = mime;
job->status.store(job->cancel.load() ? 3 : 1);
fprintf(stderr, "[Server] Job %s done (decode)\n", job->id.c_str());
}
// encode worker: VAE encode only. Encodes 48kHz interleaved stereo
// audio into latents [T_25Hz, 64] time-major, stores raw f32 in job.
// Prefers ONNX/TRT encoder when available (faster via TensorRT fusion),
// falls back to GGML encoder for GGUF/safetensors VAE models.
static void vae_encode_worker(std::shared_ptr<Job> job, AceRequest ace_req, float * src_interleaved, int src_len) {
struct buf_guard {
float * p;
~buf_guard() { if (p) free(p); }
} buf{ src_interleaved };
if (job->cancel.load()) {
job->status.store(3);
return;
}
int T_latent_max = src_len / 1920 + 64;
if (T_latent_max > MAX_T_LATENT) {
T_latent_max = MAX_T_LATENT;
}
std::vector<float> latent((size_t) T_latent_max * LATENT_CHANNELS);
int T_latent = -1;
std::string vae_name_used;
auto t_start = std::chrono::steady_clock::now();
// ── Try ONNX encoder first ─────────────────────────────────────
// Look for a *_encoder.onnx file matching the selected (or default) VAE.
// E.g., if user selected "scragvae_decoder.onnx", look for "scragvae_encoder.onnx".
// Also auto-detect from the onnx/ directory if no specific VAE is selected.
bool tried_ort = false;
{
std::string enc_onnx_path;
// If a specific VAE was requested and it's ONNX, derive encoder path
if (!ace_req.vae.empty()) {
const ModelEntry * entry = registry_find(g_registry.vae, ace_req.vae.c_str());
if (entry && entry->name.size() >= 5 &&
entry->name.substr(entry->name.size() - 5) == ".onnx") {
// Replace "_decoder.onnx" with "_encoder.onnx"
std::string p = entry->path;
auto pos = p.rfind("_decoder.onnx");
if (pos != std::string::npos) {
enc_onnx_path = p.substr(0, pos) + "_encoder.onnx";
}
}
}
// If no specific ONNX VAE selected, check the registry for any ONNX decoder
// and derive the encoder path from it
if (enc_onnx_path.empty()) {
for (const auto & e : g_registry.vae) {
if (e.name.size() >= 5 && e.name.substr(e.name.size() - 5) == ".onnx") {
std::string p = e.path;
auto pos = p.rfind("_decoder.onnx");
if (pos != std::string::npos) {
std::string candidate = p.substr(0, pos) + "_encoder.onnx";
FILE * f = fopen(candidate.c_str(), "rb");
if (f) {
fclose(f);
enc_onnx_path = candidate;
break;
}
}
}
}
}
// If we found an encoder ONNX, try ORT
if (!enc_onnx_path.empty()) {
FILE * f = fopen(enc_onnx_path.c_str(), "rb");
if (f) {
fclose(f);
tried_ort = true;
ModelKey ort_key;
ort_key.kind = MODEL_VAE_ENC_ORT;
ort_key.path = enc_onnx_path;
VaeEncOrt * enc_ort = store_require_vae_enc_ort(g_store, ort_key);
if (enc_ort) {
ModelHandle guard(g_store, enc_ort);
T_latent = vae_enc_ort_encode_tiled(enc_ort, src_interleaved, src_len,
latent.data(), T_latent_max,
g_synth_params.vae_chunk, g_synth_params.vae_overlap);
if (T_latent >= 0) {
// Extract basename for logging
auto slash = enc_onnx_path.find_last_of("/\\");
vae_name_used = (slash != std::string::npos) ? enc_onnx_path.substr(slash + 1) : enc_onnx_path;
} else {
fprintf(stderr, "[Server] encode: ORT encode failed, falling back to GGML\n");
}
} else {
fprintf(stderr, "[Server] encode: ORT session load failed, falling back to GGML\n");
}
}
}
}
// ── GGML fallback ──────────────────────────────────────────────
if (T_latent < 0) {
const ModelEntry * vae_entry = registry_find_non_onnx(g_registry.vae, ace_req.vae.c_str());
if (!vae_entry) {
vae_entry = registry_find_non_onnx(g_registry.vae);
}
if (!vae_entry) {
fprintf(stderr, "[Server] encode: no GGUF/safetensors VAE available for encoding\n");
job->status.store(2);
return;
}
ModelKey vae_key;
vae_key.kind = MODEL_VAE_ENC;
vae_key.path = vae_entry->path;
vae_key.adapter_scale = 1.0f;
VAEEncoder * vae = store_require_vae_enc(g_store, vae_key);
if (!vae) {
fprintf(stderr, "[Server] encode: store_require_vae_enc failed\n");
job->status.store(2);
return;
}
ModelHandle vae_guard(g_store, vae);
T_latent = vae_enc_encode_tiled(vae, src_interleaved, src_len, latent.data(), T_latent_max,
g_synth_params.vae_chunk, g_synth_params.vae_overlap);
if (T_latent < 0) {
fprintf(stderr, "[Server] encode: vae_enc_encode_tiled failed\n");
job->status.store(2);
return;
}
vae_name_used = vae_entry->name;
}
auto t_end = std::chrono::steady_clock::now();
float ms = (float) std::chrono::duration_cast<std::chrono::microseconds>(t_end - t_start).count() / 1000.0f;
fprintf(stderr, "[Server] encode: %d audio samples (%.2fs) -> %d latent frames, %.0fms (%s)\n", src_len,
(float) src_len / 48000.0f, T_latent, ms, vae_name_used.c_str());
if (g_keep_loaded) {
g_loaded_vae = vae_name_used;
} else {
g_loaded_vae.clear();
}
std::string body;
body.resize((size_t) T_latent * LATENT_FRAME_BYTES);
std::memcpy(body.data(), latent.data(), body.size());
job->result_body = std::move(body);
job->result_mime = "application/octet-stream";
job->status.store(job->cancel.load() ? 3 : 1);
fprintf(stderr, "[Server] Job %s done (encode)\n", job->id.c_str());
}
// POST /vae
// multipart/form-data: single VAE entrypoint, direction depends on input.
// part "audio": WAV or MP3 source audio -> encode path, latents out
// part "src_latents": raw f32 latent bytes -> decode path, audio out
// part "request": JSON text (optional, for VAE selection, output format)
// Returns JSON {"id":"N"} immediately. Result is raw latent bytes (encode)
// or audio (decode). Only one direction at a time.
static void handle_vae(const httplib::Request & req, httplib::Response & res) {
if (g_registry.vae.empty()) {
json_error(res, 501, "VAE endpoint requires a VAE in the registry");
return;
}
if (!req.is_multipart_form_data()) {
json_error(res, 400, "VAE endpoint requires multipart/form-data");
return;
}
AceRequest ace_req;
request_init(&ace_req);
if (req.form.has_file("request")) {
const std::string & json = req.form.get_file("request").content;
if (!request_parse_json(&ace_req, json.c_str())) {
json_error(res, 400, "Multipart: invalid JSON in 'request' part");
return;
}
} else if (req.form.has_field("request")) {
const std::string & json = req.form.get_field("request");
if (!request_parse_json(&ace_req, json.c_str())) {
json_error(res, 400, "Multipart: invalid JSON in 'request' part");
return;
}
}
bool has_audio = req.form.has_file("audio");
bool has_latents = req.form.has_file("src_latents");
if (has_audio == has_latents) {
json_error(res, 400, "Multipart: provide exactly one of 'audio' (encode) or 'src_latents' (decode)");
return;
}
if (has_audio) {
// encode path: audio in -> raw latents out
const auto & file = req.form.get_file("audio");
if (file.content.empty()) {
json_error(res, 400, "Multipart: empty 'audio' part");
return;
}
int T_audio = 0;
float * planar = audio_read_48k_buf((const uint8_t *) file.content.data(), file.content.size(), &T_audio);
if (!planar || T_audio <= 0) {
if (planar) free(planar);
json_error(res, 400, "Failed to decode audio");
return;
}
if ((int64_t) T_audio / 1920 >= (int64_t) MAX_T_LATENT) {
free(planar);
json_error(res, 413, "audio exceeds max duration (10 min)");
return;
}
float * src_interleaved = audio_planar_to_interleaved(planar, T_audio);
free(planar);
int src_len = T_audio;
auto job = job_create();
fprintf(stderr, "[Server] Job %s created (vae encode, %.2fs audio)\n", job->id.c_str(),
(float) src_len / 48000.0f);
work_push([job, ace_req, src_interleaved, src_len]() mutable {
vae_encode_worker(job, ace_req, src_interleaved, src_len);
});
std::string body = "{\"id\":\"" + job->id + "\"}";
res.set_content(body, "application/json");
return;
}
// decode path: raw latents in -> audio out
const auto & file = req.form.get_file("src_latents");
if (file.content.empty() || (file.content.size() % LATENT_FRAME_BYTES) != 0) {
json_error(res, 400, "src_latents size not a multiple of 64*4 bytes");
return;
}
int T = (int) (file.content.size() / (size_t) LATENT_FRAME_BYTES);
if (T > MAX_T_LATENT) {
json_error(res, 413, "src_latents exceeds max frames");
return;
}
std::vector<float> src_latents(reinterpret_cast<const float *>(file.content.data()),
reinterpret_cast<const float *>(file.content.data()) + (size_t) T * LATENT_CHANNELS);
bool output_wav = false;
WavFormat wav_fmt = WAV_S16;
{
bool is_mp3 = true;
if (!audio_parse_format(ace_req.output_format.c_str(), is_mp3, wav_fmt)) {
json_error(res, 400, "Invalid output_format (use: mp3, wav16, wav24, wav32)");
return;
}
output_wav = !is_mp3;
}
int peak_clip = ace_req.peak_clip;
auto job = job_create();
fprintf(stderr, "[Server] Job %s created (vae decode, %d latent frames)\n", job->id.c_str(), T);
work_push([job, ace_req, latents = std::move(src_latents), T, output_wav, wav_fmt, peak_clip]() mutable {
vae_decode_worker(job, ace_req, std::move(latents), T, output_wav, wav_fmt, peak_clip);
});
std::string body = "{\"id\":\"" + job->id + "\"}";
res.set_content(body, "application/json");
}
// ────────────────────────────────────────────────────────────────────────
// /codes-decode endpoint: 5 Hz LM audio codes -> audio, straight through the
// FSQ detokenizer and the VAE. Zero DiT, zero sound adapter, zero sampler —
// this renders the planner LM's plan literally, which is what the Training
// Studio's codes audition plays. docs/plans/2026-07-28-codes-preview.md §4.
// ────────────────────────────────────────────────────────────────────────
// codes-decode worker: FSQ-detokenize the codes into VAE-ready latents, then
// tail-call vae_decode_worker UNMODIFIED. detok_ggml_decode writes ggml
// [64, T_25Hz] where element (c, t) = data[t * 64 + c] (fsq-detok.h) — that is
// byte-identical to the [T, 64] time-major buffer /vae's src_latents part
// already hands vae_decode_worker, so there is no transpose and no new decode
// code. The DiT mask channel (ctx_ch = Oc*2, pipeline-synth-ops.cpp) is DiT
// conditioning applied *after* the detokenizer, so its absence here is exactly
// what "zero DiT influence" means.
static void codes_decode_worker(std::shared_ptr<Job> job,
AceRequest ace_req,
std::vector<int> codes,
std::string dit_path,
bool output_wav,
WavFormat wav_fmt,
int peak_clip) {
if (job->cancel.load()) {
job->status.store(3);
return;
}
const int T_5Hz = (int) codes.size();
const int T_25Hz = T_5Hz * 5;
// Keyed on the DiT path exactly as pipeline-synth.cpp keys its
// fsq_detok_key — the detokenizer weights live inside the DiT file.
ModelKey k;
k.kind = MODEL_FSQ_DETOK;
k.path = dit_path;
auto t_start = std::chrono::steady_clock::now();
std::vector<float> latents((size_t) T_25Hz * LATENT_CHANNELS);
int rc = -1;
{
DetokGGML * detok = store_require_fsq_detok(g_store, k);
if (!detok) {
fprintf(stderr, "[Server] codes-decode: store_require_fsq_detok failed (%s)\n", dit_path.c_str());
job->status.store(2);
return;
}
// Released before vae_decode_worker requires the VAE: under
// EVICT_STRICT a live refcount on the detokenizer would abort the store.
ModelHandle detok_guard(g_store, detok);
if (!g_synth_params.use_fa) {
detok->use_flash_attn = false;
}
rc = detok_ggml_decode(detok, codes.data(), T_5Hz, latents.data());
}
if (rc < 0) {
fprintf(stderr, "[Server] codes-decode: detok_ggml_decode failed\n");
job->status.store(2);
return;
}
auto t_end = std::chrono::steady_clock::now();
float detok_ms = (float) std::chrono::duration_cast<std::chrono::microseconds>(t_end - t_start).count() / 1000.0f;
fprintf(stderr, "[Server] codes-decode: %d codes -> %d latent frames, detok %.0f ms\n", T_5Hz, T_25Hz, detok_ms);
// The shipped VAE decode path, called not copied. Cancellation,
// g_loaded_vae stickiness, audio_normalize, peak_clip, wav/mp3 encoding and
// every job->status transition are inherited, not reimplemented.
vae_decode_worker(job, std::move(ace_req), std::move(latents), T_25Hz, output_wav, wav_fmt, peak_clip);
}
// POST /codes-decode — body is an AceRequest JSON (the same parser as /lm).
// Fields read, all others ignored:
// audio_codes REQUIRED. Comma-separated 5 Hz FSQ indices, the exact string
// /lm returns and lm_codes.jsonl stores.
// synth_model DiT registry name — the FSQ detokenizer weights live inside
// the DiT file. Empty -> resolve_name default.
// vae VAE registry name. Empty -> resolve_name default.
// output_format mp3 | wav16 | wav24 | wav32. Default mp3 (request_init).
// peak_clip int, default 10.
// Returns {"id":"N"}; poll GET /job?id=N and fetch GET /job?id=N&result=1,
// exactly like /vae's decode path — the result IS that code path's result.
static void handle_codes_decode(const httplib::Request & req, httplib::Response & res) {
if (g_registry.dit.empty() || g_registry.vae.empty()) {
json_error(res, 501, "codes-decode requires DiT and VAE models");
return;
}
AceRequest ace_req;
request_init(&ace_req);
if (!request_parse_json(&ace_req, req.body.c_str())) {
json_error(res, 400, "Invalid JSON");
return;
}
if (ace_req.audio_codes.empty()) {
json_error(res, 400, "audio_codes is required");
return;
}
// Local CSV -> int parse. parse_csv<int> is a file-local template in
// pipeline-synth-ops.cpp, not exported, so this is a local loop with the
// same separator set (',' and ' ' only — a tab/newline bails there too, and
// silently accepting them here would make a preview disagree with what the
// /synth path would have consumed from the identical string).
//
// It is deliberately STRICTER than parse_csv in one respect: parse_csv stops
// at the first non-numeric token, and this endpoint's whole product is "the
// LM's plan rendered literally". Stopping early would decode a 150-code plan
// as its first 2 codes and still return 200 with a 0.4 s WAV, while the
// caller computes duration from its own codes.split(',').length and renders
// "150 codes / 30.0s" over it. There is no field in the {"id":"N"} reply or
// in the raw-audio job result that could expose the truncation, so a
// malformed token is a 400 here, not a silent short decode.
std::vector<int> codes;
{
// Codebook cardinality from the ONE FSQ definition (fsq-quant.h, visible
// via model-store.h -> fsq-detok.h). Computed, never hardcoded, so a
// levels change cannot leave a stale 64000 behind.
int64_t fsq_codebook = 1;
for (int d = 0; d < FSQ_NDIMS; d++) {
fsq_codebook *= (int64_t) FSQ_LEVELS[d];
}
size_t out_of_codebook = 0;
long long first_bad_code = LLONG_MIN;
const char * base = ace_req.audio_codes.c_str();
const char * p = base;
const char * end = p + ace_req.audio_codes.size();
while (p < end) {
while (p < end && (*p == ',' || *p == ' ')) {
p++;
}
if (p >= end) {
break;
}
char * next = nullptr;
long long v = strtoll(p, &next, 10);
if (next == p) {
std::string msg = "audio_codes has a malformed token at offset " + std::to_string(p - base) +
" (expected comma-separated base-10 integers)";
json_error(res, 400, msg.c_str());
return;
}
// Overflow is unambiguously malformed — no LM emits this, and the
// (int) cast below would otherwise silently truncate it to an
// implementation-defined, possibly negative value.
if (v < INT32_MIN || v > INT32_MAX) {
std::string msg = "audio_codes has an out-of-range token " + std::to_string(v) + " at offset " +
std::to_string(p - base);
json_error(res, 400, msg.c_str());
return;
}
// Out-of-CODEBOOK is only WARNED about, deliberately not rejected.
// fsq_decode_index does (index / stride) % L with no clamp, so a code
// >= the codebook wraps to a different entry and a negative one puts
// out[d] below -1 — both feed out-of-distribution latents to the
// detokenizer. Tempting to 400. But the LM's audio-code band is
// AUDIO_CODE_COUNT = 65535 wide (prompt.h) while the codebook is only
// 64000, and metadata-fsm.h's mask spans the FULL band — so codes
// 64000..65534 are samplable and may well appear in legitimate output.
// Rejecting them would fail a real audition outright, which is far
// worse than the wrapped audio it prevents. This stays a log line
// (enough to identify a wrong-vocabulary adapter) until a measurement
// shows whether real LM output ever exceeds the codebook.
if (v < 0 || v >= fsq_codebook) {
out_of_codebook++;
if (first_bad_code == LLONG_MIN) {
first_bad_code = v;
}
}
codes.push_back((int) v);
p = next;
}
if (out_of_codebook > 0) {
fprintf(stderr,
"[Server] codes-decode: WARNING %zu/%zu codes are outside the FSQ codebook 0..%lld "
"(first: %lld) — these wrap in fsq_decode_index and will decode as noise; a "
"wrong-vocabulary LM adapter is the usual cause\n",
out_of_codebook, codes.size(), (long long) (fsq_codebook - 1), first_bad_code);
}
}
if (codes.empty()) {
json_error(res, 400, "audio_codes contains no valid codes");
return;
}
if ((int64_t) codes.size() * 5 > (int64_t) MAX_T_LATENT) {
json_error(res, 413, "audio_codes exceeds max duration (10 min)");
return;
}
bool output_wav = false;
WavFormat wav_fmt = WAV_S16;
{
bool is_mp3 = true;
if (!audio_parse_format(ace_req.output_format.c_str(), is_mp3, wav_fmt)) {
json_error(res, 400, "Invalid output_format (use: mp3, wav16, wav24, wav32)");
return;
}
output_wav = !is_mp3;
}
int peak_clip = ace_req.peak_clip;
// Resolve the DiT here, not in the worker, so an ONNX DiT is a 400 with an
// explanation instead of a silent job failure: the FSQ detokenizer weights
// are not part of an ONNX export. dit_ends_with_onnx comes from dit.h,
// already visible via model-store.h — no new include in this TU.
//
// The sticky hint is deliberately EMPTY (not g_loaded_dit). synth_worker
// overwrites g_loaded_dit on every /synth completion under EVICT_NEVER —
// and the audition flow posts /lm?keep_loaded=1, which latches EVICT_NEVER
// for the process lifetime, so that global is very much live here. With the
// sticky hint in play, two identical requests carrying synth_model:"" either
// side of one Create-panel generation would decode through two different FSQ
// detokenizers and produce different audio — exactly what an A/B audition
// cannot tolerate, and what a back-to-back determinism check cannot detect.
// Empty hint => registry[0], which is stable for the process lifetime.
std::string dit_name = resolve_name(g_registry.dit, ace_req.synth_model, std::string());
const ModelEntry * dit = registry_find(g_registry.dit, dit_name.c_str());
if (!dit) {
std::string msg = "DiT not found: " + dit_name;
json_error(res, 400, msg.c_str());
return;
}
if (dit_ends_with_onnx(dit->path.c_str())) {
json_error(res, 400,
"codes-decode needs a GGUF or safetensors DiT — the FSQ detokenizer weights are not in an "
"ONNX export");
return;
}
std::string dit_path = dit->path;
// Same treatment for the VAE, for the same two reasons. vae_decode_worker
// resolves with plain registry_find against the sticky g_loaded_vae, so
// (a) an ONNX VAE — registry_scan does put those in reg->vae — would reach
// store_require_vae_dec and fail the job with nothing but a log line, which
// is precisely the opaque failure the DiT check three lines up exists to
// prevent, and (b) an empty `vae` would inherit whatever the last /synth
// latched. Resolving to a non-ONNX entry HERE and pinning the resolved name
// back into the request makes the worker's resolve_name return it verbatim
// (a non-empty request field always wins) without modifying the worker.
// registry_find_non_onnx(bucket, nullptr) already means "first non-ONNX".
const ModelEntry * vae = registry_find_non_onnx(g_registry.vae, ace_req.vae.empty() ? nullptr : ace_req.vae.c_str());
if (!vae) {
std::string msg = ace_req.vae.empty()
? std::string("codes-decode needs a GGUF or safetensors VAE — no non-ONNX VAE is "
"registered")
: ("VAE not found (or is an ONNX export, which has no usable decoder here): " +
ace_req.vae);
json_error(res, 400, msg.c_str());
return;
}
ace_req.vae = vae->name;
auto job = job_create();
fprintf(stderr, "[Server] Job %s created (codes-decode, %zu codes, dit=%s)\n", job->id.c_str(), codes.size(),
dit_name.c_str());
work_push([job, ace_req, codes = std::move(codes), dit_path, output_wav, wav_fmt, peak_clip]() mutable {
codes_decode_worker(job, std::move(ace_req), std::move(codes), dit_path, output_wav, wav_fmt, peak_clip);
});
std::string body = "{\"id\":\"" + job->id + "\"}";
res.set_content(body, "application/json");
}
// warm worker: same setup as synth_worker through ace_synth_load, then stops.
// Under EVICT_NEVER (set by `--keep-loaded` or a prior `?keep_loaded=1`) the
// modules stay resident, so the next /synth using the same DiT + adapter combo
// skips the cold-start load. Returns immediately on STRICT — pre-loading there
// is pointless since the modules would be evicted instantly.
struct WarmRequest {
std::string dit;
std::string vae;
std::string adapter;
float adapter_scale = 1.0f;
};
static void warm_worker(std::shared_ptr<Job> job, WarmRequest wr) {
if (job->cancel.load()) {
job_set_phase(*job, JobPhase::CANCELLED);
job->status.store(3);
return;
}
if (!g_keep_loaded) {
// Document the no-op clearly so callers learn from the response.
job->result_body = "{\"warm\":false,\"reason\":\"keep_loaded not set; would be evicted immediately\"}";
job->result_mime = "application/json";
job_set_phase(*job, JobPhase::DONE);
job->status.store(1);
return;
}
std::string dit_name = resolve_name(g_registry.dit, wr.dit, g_loaded_dit);
const ModelEntry * dit = registry_find(g_registry.dit, dit_name.c_str());
if (!dit || g_registry.text_enc.empty() || g_registry.vae.empty()) {
fprintf(stderr, "[Server] warm: DiT/Text-Enc/VAE not resolvable\n");
job_set_phase(*job, JobPhase::FAILED);
job->status.store(2);
return;
}
const ModelEntry * vae = registry_find_non_onnx(g_registry.vae);
if (!wr.vae.empty()) {
const ModelEntry * sel = registry_find(g_registry.vae, wr.vae.c_str());
if (sel && !(sel->name.size() >= 5 && sel->name.substr(sel->name.size() - 5) == ".onnx")) {
vae = sel;
}
}
if (!vae) {
fprintf(stderr, "[Server] warm: no non-ONNX VAE available\n");
job_set_phase(*job, JobPhase::FAILED);
job->status.store(2);
return;
}
std::string vae_name = vae->name;
// Idempotent fast-path: if the same DiT + adapter + VAE combo is already in
// the sticky-name slot, the store holds it resident — nothing to do.
if (g_loaded_dit == dit_name && g_loaded_adapter == wr.adapter && g_loaded_vae == vae_name) {
job->result_body = "{\"warm\":true,\"already_loaded\":true}";
job->result_mime = "application/json";
job_set_phase(*job, JobPhase::DONE);
job->status.store(1);
fprintf(stderr, "[Server] warm: already loaded (DiT=%s VAE=%s Adapter=%s)\n",
dit_name.c_str(), vae_name.c_str(), wr.adapter.c_str());
return;
}
AceSynthParams p = g_synth_params;
p.text_encoder_path = g_registry.text_enc[0].path.c_str();
p.dit_path = dit->path.c_str();
p.vae_path = vae->path.c_str();
p.adapter_path = nullptr;
p.adapter_scale = 1.0f;
// Mirror the synth worker's stack so warm and real loads produce the same
// DiT cache key (the warm endpoint takes a single adapter; fold it into a
// one-element stack). Clearing first prevents a stale stack from a prior
// synth request leaking into this load.
g_hotstep_params.adapters.clear();
if (!wr.adapter.empty()) {
const AdapterEntry * adapter = registry_find_adapter(g_registry, wr.adapter.c_str());
if (!adapter) {
fprintf(stderr, "[Server] warm: adapter not found: %s\n", wr.adapter.c_str());
job_set_phase(*job, JobPhase::FAILED);
job->status.store(2);
return;
}
g_hotstep_params.adapters.push_back({ adapter->path, wr.adapter_scale });
p.adapter_path = g_hotstep_params.adapters[0].path.c_str();
p.adapter_scale = g_hotstep_params.adapters[0].scale;
}
fprintf(stderr, "[Server] warm: loading DiT=%s VAE=%s%s%s\n", dit_name.c_str(), vae_name.c_str(),
wr.adapter.empty() ? "" : " Adapter=", wr.adapter.c_str());
job_set_phase(*job, wr.adapter.empty() ? JobPhase::LOADING_DIT : JobPhase::ADAPTER_PRECOMPUTE);
// Wire the per-job cancel flag into the adapter precompute loops so a cancel
// during cold start aborts in <100 ms. Cleared via RAII on every exit path.
g_adapter_cancel.store(&job->cancel, std::memory_order_release);
struct WarmCancelGuard {
~WarmCancelGuard() { g_adapter_cancel.store(nullptr, std::memory_order_release); }
} warm_cancel_guard;
AceSynth * ctx = ace_synth_load(g_store, &p);
if (!ctx) {
fprintf(stderr, "[Server] warm: synth load failed\n");
bool cancelled = job->cancel.load();
job_set_phase(*job, cancelled ? JobPhase::CANCELLED : JobPhase::FAILED);
job->status.store(cancelled ? 3 : 2);
return;
}
// Free the ctx but leave the underlying store entries resident — under
// EVICT_NEVER the store ignores the refcount drop, so the modules stay hot.
ace_synth_free(ctx);
// Set the sticky-name hints so resolve_name picks the same models next.
g_loaded_dit = dit_name;
g_loaded_adapter = wr.adapter;
g_loaded_adapter_scale = wr.adapter_scale;
g_loaded_vae = vae_name;
job->result_body = "{\"warm\":true,\"already_loaded\":false}";
job->result_mime = "application/json";
job_set_phase(*job, JobPhase::DONE);
job->status.store(1);
fprintf(stderr, "[Server] warm: done (DiT=%s VAE=%s Adapter=%s)\n", dit_name.c_str(), vae_name.c_str(),
wr.adapter.c_str());
}
// POST /warm
// Body { dit, vae?, adapter?, adapter_scale? }. Spawns a background job that
// loads the requested DiT + VAE + adapter so the next /synth with the same key
// short-circuits the model load. Honors `?keep_loaded=1`. Returns {"id":"N"};
// poll via GET /job?id=N as with /synth.
static void handle_warm(const httplib::Request & req, httplib::Response & res) {
if (g_registry.dit.empty() || g_registry.text_enc.empty() || g_registry.vae.empty()) {
json_error(res, 501, "No synth models in registry (need dit + text-encoder + vae)");
return;
}
// per-request co-resident: ?keep_loaded=1 flips the store to NEVER for the
// rest of the process lifetime (same one-way behavior as /synth and /lm).
const bool req_keep_loaded = req.has_param("keep_loaded") && req.get_param_value("keep_loaded") == "1";
if (req_keep_loaded && !g_keep_loaded) {
g_keep_loaded = true;
store_set_policy(g_store, EVICT_NEVER);
fprintf(stderr, "[Server] keep_loaded enabled via /warm request query param\n");
}
WarmRequest wr;
if (!req.body.empty()) {
yyjson_doc * doc = yyjson_read(req.body.c_str(), req.body.size(), 0);
yyjson_val * root = doc ? yyjson_doc_get_root(doc) : nullptr;
if (!root || !yyjson_is_obj(root)) {
if (doc) yyjson_doc_free(doc);
json_error(res, 400, "Invalid JSON: expected an object");
return;
}
yyjson_val * v;
if ((v = yyjson_obj_get(root, "dit")) && yyjson_is_str(v)) wr.dit = yyjson_get_str(v);
if ((v = yyjson_obj_get(root, "vae")) && yyjson_is_str(v)) wr.vae = yyjson_get_str(v);
if ((v = yyjson_obj_get(root, "adapter")) && yyjson_is_str(v)) wr.adapter = yyjson_get_str(v);
if ((v = yyjson_obj_get(root, "adapter_scale")) && yyjson_is_num(v)) {
wr.adapter_scale = (float) yyjson_get_num(v);
}
yyjson_doc_free(doc);
}
auto job = job_create();
fprintf(stderr, "[Server] Job %s created (warm: dit=%s vae=%s adapter=%s)\n", job->id.c_str(), wr.dit.c_str(),
wr.vae.c_str(), wr.adapter.c_str());
work_push([job, wr]() mutable { warm_worker(job, std::move(wr)); });
std::string body = "{\"id\":\"" + job->id + "\"}";
res.set_content(body, "application/json");
}
// GET /jobs
// Array of every job currently in g_jobs. Lets external reconcilers discover
// live engine jobs when a client died mid-poll. Honors the existing MAX_JOBS
// eviction policy. Read-only — does not touch the worker queue or model store.
static void handle_jobs_list(const httplib::Request &, httplib::Response & res) {
yyjson_mut_doc * doc = yyjson_mut_doc_new(NULL);
yyjson_mut_val * arr = yyjson_mut_arr(doc);
yyjson_mut_doc_set_root(doc, arr);
std::lock_guard<std::mutex> lock(mtx_jobs);
for (const auto & id : g_job_order) {
auto it = g_jobs.find(id);
if (it == g_jobs.end()) continue;
const auto & j = it->second;
yyjson_mut_val * obj = yyjson_mut_obj(doc);
yyjson_mut_obj_add_str(doc, obj, "id", j->id.c_str());
yyjson_mut_obj_add_str(doc, obj, "status", job_status_str(j->status.load()));
yyjson_mut_obj_add_str(doc, obj, "phase", job_phase_str(j->phase.load()));
yyjson_mut_obj_add_int(doc, obj, "phase_step", j->phase_step.load(std::memory_order_relaxed));
yyjson_mut_obj_add_int(doc, obj, "phase_total", j->phase_total.load(std::memory_order_relaxed));
yyjson_mut_arr_append(arr, obj);
}
char * json = yyjson_mut_write(doc, 0, NULL);
yyjson_mut_doc_free(doc);
res.set_content(json ? json : "[]", "application/json");
if (json) free(json);
}
// GET /props
// server configuration, available models, and default request.
// the webui reads this at boot to populate dropdowns and status indicators.
static void handle_props(const httplib::Request &, httplib::Response & res) {
yyjson_mut_doc * doc = yyjson_mut_doc_new(NULL);
yyjson_mut_val * root = yyjson_mut_obj(doc);
yyjson_mut_doc_set_root(doc, root);
yyjson_mut_obj_add_str(doc, root, "version", ACE_VERSION);
// helper: build a JSON array of model entry names
auto add_names = [&](yyjson_mut_val * parent, const char * key, const std::vector<ModelEntry> & bucket) {
yyjson_mut_val * arr = yyjson_mut_arr(doc);
for (const auto & e : bucket) {
yyjson_mut_arr_add_str(doc, arr, e.name.c_str());
}
yyjson_mut_obj_add_val(doc, parent, key, arr);
};
// models: available model names per bucket
yyjson_mut_val * models = yyjson_mut_obj(doc);
yyjson_mut_obj_add_val(doc, root, "models", models);
add_names(models, "lm", g_registry.lm);
add_names(models, "embedding", g_registry.text_enc);
add_names(models, "dit", g_registry.dit);
add_names(models, "vae", g_registry.vae);
// adapters: available adapter names
yyjson_mut_val * adapters_arr = yyjson_mut_arr(doc);
for (const auto & e : g_registry.adapters) {
yyjson_mut_arr_add_str(doc, adapters_arr, e.name.c_str());
}
yyjson_mut_obj_add_val(doc, root, "adapters", adapters_arr);
// lm_adapters: planner-LM LoRAs (local HOT-Step feature, adapters/lm/)
yyjson_mut_val * lm_adapters_arr = yyjson_mut_arr(doc);
for (const auto & e : g_registry.lm_adapters) {
yyjson_mut_arr_add_str(doc, lm_adapters_arr, e.name.c_str());
}
yyjson_mut_obj_add_val(doc, root, "lm_adapters", lm_adapters_arr);
// cli: server settings
yyjson_mut_val * cli = yyjson_mut_obj(doc);
yyjson_mut_obj_add_val(doc, root, "cli", cli);
yyjson_mut_obj_add_int(doc, cli, "max_batch", g_max_batch);
yyjson_mut_obj_add_int(doc, cli, "mp3_bitrate", g_mp3_kbps);
// default: full AceRequest with all defaults from request_init().
// the webui reads this to populate LM placeholders.
// DiT fields (inference_steps, guidance_scale, shift) are 0 = auto-detect;
// their resolved placeholders come from presets below.
AceRequest defaults;
request_init(&defaults);
std::string defaults_str = request_to_json(&defaults, false);
yyjson_doc * defaults_doc = yyjson_read(defaults_str.c_str(), defaults_str.size(), 0);
yyjson_mut_val * defaults_copy = yyjson_val_mut_copy(doc, yyjson_doc_get_root(defaults_doc));
yyjson_mut_obj_add_val(doc, root, "default", defaults_copy);
yyjson_doc_free(defaults_doc);
// presets: auto-detect values for DiT sampling params.
// the webui switches placeholders based on the selected DiT model.
yyjson_mut_val * presets = yyjson_mut_obj(doc);
yyjson_mut_obj_add_val(doc, root, "presets", presets);
yyjson_mut_val * turbo = yyjson_mut_obj(doc);
yyjson_mut_obj_add_int(doc, turbo, "inference_steps", 8);
yyjson_mut_obj_add_real(doc, turbo, "guidance_scale", 1.0);
yyjson_mut_obj_add_real(doc, turbo, "shift", 3.0);
yyjson_mut_obj_add_val(doc, presets, "turbo", turbo);
yyjson_mut_val * sft = yyjson_mut_obj(doc);
yyjson_mut_obj_add_int(doc, sft, "inference_steps", 50);
yyjson_mut_obj_add_real(doc, sft, "guidance_scale", 1.0);
yyjson_mut_obj_add_real(doc, sft, "shift", 1.0);
yyjson_mut_obj_add_val(doc, presets, "sft", sft);
// serialize
yyjson_write_flag flags = YYJSON_WRITE_PRETTY | YYJSON_WRITE_PRETTY_TWO_SPACES | YYJSON_WRITE_FP_TO_FIXED(2);
char * json = yyjson_mut_write(doc, flags, NULL);
yyjson_mut_doc_free(doc);
res.set_content(json, "application/json");
free(json);
}
static void usage(const char * prog) {
AceLmParams lm_d;
AceSynthParams synth_d;
ace_lm_default_params(&lm_d);
ace_synth_default_params(&synth_d);
fprintf(stderr, "acestep.cpp %s\n\n", ACE_VERSION);
fprintf(stderr,
"Usage: %s --models <dir> [options]\n"
"\n"
"Required:\n"
" --models <dir> Directory of GGUF model files\n"
"\n"
"Adapter:\n"
" --adapters <dir> Directory of adapters\n"
"\n"
"Memory control:\n"
" --keep-loaded Keep models in VRAM between requests\n"
" --vae-chunk <N> Latent frames per tile (default: %d)\n"
" --vae-overlap <N> Overlap frames per side (default: %d)\n"
"\n"
"ONNX/TensorRT:\n"
" --onnx-dir <dir> Directory with ONNX models (e.g. vae_decoder.onnx)\n"
"\n"
"Speculative decoding:\n"
" --draft-lm <path> Path to 0.6B draft LM (auto-discovers if omitted)\n"
" --no-draft Disable draft model auto-discovery\n"
"\n"
"Output:\n"
" --mp3-bitrate <kbps> MP3 bitrate (default: %d)\n"
"\n"
"Server:\n"
" --host <addr> Listen address (default: 127.0.0.1)\n"
" --port <N> Listen port (default: 8080)\n"
" --max-batch <N> LM batch limit (default: %d)\n"
" --max-seq <N> KV cache size (default: %d)\n"
"\n"
"Debug:\n"
" --no-fsm Disable FSM constrained decoding\n"
" --no-fa Disable flash attention\n"
" --no-batch-cfg Split CFG into two separate forwards (LM + DiT)\n"
" --clamp-fp16 Clamp hidden states to FP16 range\n",
prog, synth_d.vae_chunk, synth_d.vae_overlap, g_mp3_kbps, g_max_batch, lm_d.max_seq);
}
int main(int argc, char ** argv) {
ace_lm_default_params(&g_lm_params);
ace_synth_default_params(&g_synth_params);
const char * host = "127.0.0.1";
int port = 8080;
const char * models_dir = nullptr;
const char * adapters_dir = nullptr;
const char * noise_profile_path = nullptr;
if (argc < 2) {
usage(argv[0]);
return 1;
}
for (int i = 1; i < argc; i++) {
if (!strcmp(argv[i], "--models") && i + 1 < argc) {
models_dir = argv[++i];
} else if (!strcmp(argv[i], "--adapters") && i + 1 < argc) {
adapters_dir = argv[++i];
} else if (!strcmp(argv[i], "--noise-profile") && i + 1 < argc) {
noise_profile_path = argv[++i];
} else if (!strcmp(argv[i], "--max-seq") && i + 1 < argc) {
g_lm_params.max_seq = atoi(argv[++i]);
// vae tiling
} else if (!strcmp(argv[i], "--vae-chunk") && i + 1 < argc) {
g_synth_params.vae_chunk = atoi(argv[++i]);
} else if (!strcmp(argv[i], "--vae-overlap") && i + 1 < argc) {
g_synth_params.vae_overlap = atoi(argv[++i]);
} else if (!strcmp(argv[i], "--keep-loaded")) {
g_keep_loaded = true;
g_keep_loaded_cli = true;
// output
} else if (!strcmp(argv[i], "--mp3-bitrate") && i + 1 < argc) {
g_mp3_kbps = atoi(argv[++i]);
// server
} else if (!strcmp(argv[i], "--host") && i + 1 < argc) {
host = argv[++i];
} else if (!strcmp(argv[i], "--port") && i + 1 < argc) {
port = atoi(argv[++i]);
} else if (!strcmp(argv[i], "--max-batch") && i + 1 < argc) {
g_max_batch = atoi(argv[++i]);
// debug
} else if (!strcmp(argv[i], "--no-fsm")) {
g_lm_params.use_fsm = false;
} else if (!strcmp(argv[i], "--no-fa")) {
g_lm_params.use_fa = false;
g_synth_params.use_fa = false;
} else if (!strcmp(argv[i], "--no-batch-cfg")) {
g_lm_params.use_batch_cfg = false;
g_synth_params.use_batch_cfg = false;
} else if (!strcmp(argv[i], "--clamp-fp16")) {
g_lm_params.clamp_fp16 = true;
g_synth_params.clamp_fp16 = true;
// speculative decoding
} else if (!strcmp(argv[i], "--draft-lm") && i + 1 < argc) {
g_draft_lm_path = argv[++i];
} else if (!strcmp(argv[i], "--no-draft")) {
g_draft_lm_path = "none";
} else if (!strcmp(argv[i], "--onnx-dir") && i + 1 < argc) {
g_onnx_dir = argv[++i];
} else if (!strcmp(argv[i], "--help") || !strcmp(argv[i], "-h")) {
usage(argv[0]);
return 0;
} else {
fprintf(stderr, "Unknown option: %s\n", argv[i]);
usage(argv[0]);
return 1;
}
}
// --models is required
if (!models_dir) {
fprintf(stderr, "[Server] ERROR: --models is required\n");
usage(argv[0]);
return 1;
}
// stderr capture for SSE /logs (must be after arg parsing so --help prints directly)
LogCapture log_capture;
// scan models directory (reads GGUF metadata only)
fprintf(stderr, "[Server] Scanning models in %s\n", models_dir);
if (!registry_scan(&g_registry, models_dir)) {
fprintf(stderr, "[Server] ERROR: no models found in %s\n", models_dir);
return 1;
}
// Also scan the onnx/ subdirectory for ONNX models (TRT acceleration)
{
std::string onnx_subdir = std::string(models_dir) + REGISTRY_SEP + "onnx";
registry_scan(&g_registry, onnx_subdir.c_str());
}
// speculative decoding: only via explicit --draft-lm flag
// Auto-discovery DISABLED — GGML per-call overhead (~10ms) makes the 0.6B
// draft nearly as expensive as the 4B target. Re-enable when persistent
// graphs or CUDA graph capture reduce overhead below ~2ms.
if (g_draft_lm_path == "none") {
fprintf(stderr, "[Server] Draft LM disabled (--no-draft)\n");
g_draft_lm_path.clear();
} else if (!g_draft_lm_path.empty()) {
fprintf(stderr, "[Server] Draft LM (explicit): %s\n", g_draft_lm_path.c_str());
}
// scan adapters directory (optional)
if (adapters_dir) {
fprintf(stderr, "[Server] Scanning adapters in %s\n", adapters_dir);
registry_scan_adapters(&g_registry, adapters_dir);
registry_scan_lm_adapters(&g_registry, adapters_dir);
}
// HOT-Step: load noise profile for spectral denoiser (optional)
if (noise_profile_path) {
fprintf(stderr, "[Server] Loading noise profile: %s\n", noise_profile_path);
int np_T = 0;
int np_sr = 0;
float * np_audio = audio_io_read_wav(noise_profile_path, &np_T, &np_sr);
if (np_audio && np_T > 0) {
// audio_io_read_wav returns planar stereo [L: T][R: T] — average to mono
std::vector<float> mono(np_T);
for (int i = 0; i < np_T; i++) {
mono[i] = (np_audio[i] + np_audio[np_T + i]) * 0.5f;
}
free(np_audio);
if (audio_denoise_compute_profile(mono.data(), np_T, np_sr, &g_noise_profile) == 0) {
fprintf(stderr, "[Server] Noise profile loaded successfully (%d frames, %d Hz)\n",
g_noise_profile.n_frames, g_noise_profile.sample_rate);
} else {
fprintf(stderr, "[Server] WARNING: failed to compute noise profile\n");
}
} else {
fprintf(stderr, "[Server] WARNING: could not read noise profile WAV: %s\n", noise_profile_path);
}
}
// ONNX/TensorRT: auto-detect vae_decoder.onnx in --onnx-dir
// Try new subdirectory layout first (onnx/vae/), fall back to legacy flat layout.
static std::string g_onnx_vae_path_buf;
if (g_onnx_dir) {
// Try new location: onnx_dir/vae/vae_decoder.onnx
g_onnx_vae_path_buf = std::string(g_onnx_dir) + "/vae/vae_decoder.onnx";
FILE * f = fopen(g_onnx_vae_path_buf.c_str(), "rb");
if (!f) {
// Fall back to legacy flat layout: onnx_dir/vae_decoder.onnx
g_onnx_vae_path_buf = std::string(g_onnx_dir) + "/vae_decoder.onnx";
f = fopen(g_onnx_vae_path_buf.c_str(), "rb");
}
if (f) {
fclose(f);
g_synth_params.onnx_vae_path = g_onnx_vae_path_buf.c_str();
fprintf(stderr, "[Server] ONNX VAE decoder: %s\n", g_onnx_vae_path_buf.c_str());
} else {
fprintf(stderr, "[Server] WARNING: --onnx-dir specified but no vae_decoder.onnx found in %s\n",
g_onnx_dir);
g_onnx_vae_path_buf.clear();
}
}
// validate pipeline
bool have_lm = !g_registry.lm.empty();
bool have_dit = !g_registry.dit.empty();
bool have_enc = !g_registry.text_enc.empty();
bool have_vae = !g_registry.vae.empty();
bool have_synth = have_dit && have_enc && have_vae;
// partial synth: some components found but pipeline incomplete
if (!have_synth && (have_dit || have_enc || have_vae)) {
char missing[64];
int n = 0;
if (!have_dit) {
n += snprintf(missing + n, sizeof(missing) - n, "%sDiT", n ? ", " : "");
}
if (!have_enc) {
n += snprintf(missing + n, sizeof(missing) - n, "%sText-Enc", n ? ", " : "");
}
if (!have_vae) {
n += snprintf(missing + n, sizeof(missing) - n, "%sVAE", n ? ", " : "");
}
if (have_lm) {
fprintf(stderr, "[Server] WARNING: /synth unavailable, missing: %s\n", missing);
} else {
fprintf(stderr, "[Server] ERROR: no usable pipeline, synth missing: %s\n", missing);
return 1;
}
}
// clamp max_batch
if (g_max_batch < 1) {
g_max_batch = 1;
}
if (g_max_batch > 9) {
g_max_batch = 9;
}
g_lm_params.max_batch = g_max_batch;
if (!g_draft_lm_path.empty()) {
g_lm_params.draft_model_path = g_draft_lm_path.c_str();
}
// init understand params (vae for audio encoding, dit resolved per-request)
ace_understand_default_params(&g_und_params);
g_und_params.use_fa = g_lm_params.use_fa;
g_und_params.use_fsm = g_lm_params.use_fsm;
g_und_params.max_seq = g_lm_params.max_seq; // must match ace_lm: part of the LM ModelKey
g_und_params.max_batch = g_lm_params.max_batch; // must match ace_lm: part of the LM ModelKey
g_und_params.vae_chunk = g_synth_params.vae_chunk; // share --vae-chunk with /synth
g_und_params.vae_overlap = g_synth_params.vae_overlap; // share --vae-overlap with /synth
if (have_vae) {
g_und_params.vae_path = g_registry.vae[0].path.c_str();
}
bool have_understand = have_lm && have_dit && have_vae;
// central store: one policy for the whole server lifetime. STRICT keeps
// at most one GPU module resident at a time; --keep-loaded flips it to
// NEVER and lets the working set accumulate across requests.
g_store = store_create(g_keep_loaded ? EVICT_NEVER : EVICT_STRICT);
// Initialize Lua plugin system.
// engine_dir: derive from executable path.
// Binary location varies by build system:
// - Visual Studio (multi-config): engine/build/Release/ace-server.exe (3 levels up)
// - Ninja / Makefiles / macOS: engine/build/ace-server (2 levels up)
// - Portable release: engine/ace-server (1 level up)
// Scans both engine/plugins/ (native) and <project-root>/plugins/ (community)
{
std::filesystem::path exe_path = std::filesystem::canonical(argv[0]);
std::filesystem::path exe_dir = exe_path.parent_path();
std::string dir_name = exe_dir.filename().string();
std::filesystem::path engine_dir;
if (dir_name == "Release" || dir_name == "Debug" ||
dir_name == "RelWithDebInfo" || dir_name == "MinSizeRel") {
// Multi-config generator: engine/build/Release/ → engine/ is 3 levels
engine_dir = exe_dir.parent_path().parent_path();
} else if (dir_name == "build") {
// Single-config generator: engine/build/ → engine/ is 1 level
engine_dir = exe_dir.parent_path();
} else {
// Portable release: engine/ → engine/ is 0 levels (already there)
engine_dir = exe_dir;
}
// Project root is one more level up from engine/
std::filesystem::path project_dir = engine_dir.parent_path();
PluginRegistry::instance().init(engine_dir.string(), project_dir.string());
}
// setup HTTP server
httplib::Server svr;
g_svr = &svr;
// per-operation socket idle timeout (httplib default is 5s).
// generous margin for slow networks and large audio transfers.
svr.set_read_timeout(600);
svr.set_write_timeout(600);
// SO_REUSEADDR: allow rebind after TIME_WAIT (normal restart).
// no SO_REUSEPORT: fail if another process is actively listening.
svr.set_socket_options([](socket_t sock) {
int one = 1;
#ifdef _WIN32
setsockopt(sock, SOL_SOCKET, SO_REUSEADDR, (const char *) &one, sizeof(one));
#else
setsockopt(sock, SOL_SOCKET, SO_REUSEADDR, &one, sizeof(one));
#endif
});
// reject oversized bodies (256 MB: src + ref audio, up to 10min WAV each)
svr.set_payload_max_length(256 * 1024 * 1024);
// all endpoints are always registered. handlers return 501 when the
// backing pipeline has no models in the registry.
svr.Post("/lm", handle_lm);
svr.Post("/synth", handle_synth);
svr.Post("/understand", handle_understand);
svr.Post("/vae", handle_vae);
svr.Post("/codes-decode", handle_codes_decode);
svr.Post("/warm", handle_warm);
svr.Get("/health", [](const httplib::Request &, httplib::Response & res) {
res.set_content("{\"status\":\"ok\"}", "application/json");
});
svr.Get("/props", handle_props);
svr.Get("/logs", handle_logs);
svr.Get("/jobs", handle_jobs_list);
// HOT-STEP: Lua plugin registry endpoint
svr.Get("/plugins", [](const httplib::Request &, httplib::Response & res) {
std::string json = PluginRegistry::instance().to_json();
res.set_content(json, "application/json");
});
// HOT-STEP: GET /vram — GPU memory usage (CUDA only)
svr.Get("/vram", [](const httplib::Request &, httplib::Response & res) {
#ifdef GGML_USE_CUDA
size_t free_bytes = 0, total_bytes = 0;
cudaError_t err = cudaMemGetInfo(&free_bytes, &total_bytes);
if (err != cudaSuccess) {
json_error(res, 500, cudaGetErrorString(err));
return;
}
size_t used_bytes = total_bytes - free_bytes;
char buf[256];
snprintf(buf, sizeof(buf),
"{\"used_mb\":%.0f,\"total_mb\":%.0f,\"free_mb\":%.0f}",
(double) used_bytes / (1024.0 * 1024.0),
(double) total_bytes / (1024.0 * 1024.0),
(double) free_bytes / (1024.0 * 1024.0));
res.set_content(buf, "application/json");
#else
res.set_content("{\"used_mb\":0,\"total_mb\":0,\"free_mb\":0}", "application/json");
#endif
});
// List currently-resident GPU modules (manual-unload UI).
svr.Get("/models/loaded", [](const httplib::Request &, httplib::Response & res) {
struct Acc { std::string json; bool first = true; } acc;
store_list_loaded(g_store,
[](const char * label, size_t bytes, int refcount, void * ud) {
Acc * a = static_cast<Acc *>(ud);
char buf[256];
snprintf(buf, sizeof(buf), "%s{\"label\":\"%s\",\"mb\":%.0f,\"in_use\":%s}",
a->first ? "" : ",", label, (double) bytes / (1024.0 * 1024.0),
refcount > 0 ? "true" : "false");
a->json += buf;
a->first = false;
},
&acc);
res.set_content("{\"loaded\":[" + acc.json + "]}", "application/json");
});
// Manually unload one resident module by label. In-use modules are skipped;
// under keep-loaded the module reloads on next use, so this is safe anytime.
svr.Post("/models/unload", [](const httplib::Request & req, httplib::Response & res) {
yyjson_doc * doc = yyjson_read(req.body.c_str(), req.body.size(), 0);
const char * label = nullptr;
if (doc) {
yyjson_val * root = yyjson_doc_get_root(doc);
yyjson_val * v = root ? yyjson_obj_get(root, "label") : nullptr;
if (v && yyjson_is_str(v)) label = yyjson_get_str(v);
}
if (!label || !*label) {
if (doc) yyjson_doc_free(doc);
json_error(res, 400, "missing 'label'");
return;
}
bool freed = store_evict_label(g_store, label);
char buf[160];
snprintf(buf, sizeof(buf), "{\"unloaded\":%s,\"label\":\"%s\"}", freed ? "true" : "false", label);
res.set_content(buf, "application/json");
if (doc) yyjson_doc_free(doc);
});
// HOT-STEP: undo a ?keep_loaded=1 latch. The /lm handler's comment says
// going back to STRICT "would need a full eviction pass and is not safe
// mid-flight" — this endpoint IS that pass: evict every unreferenced GPU
// module, and only when nothing stays resident (nothing mid-flight) flip
// the policy back to EVICT_STRICT and clear the latch. The codes audition
// calls this on every job exit so a single audition no longer leaves the
// engine hoarding the whole pipeline for the rest of the session.
// Refused when --keep-loaded came from the command line — that residency
// is the user's explicit choice, not job plumbing.
svr.Post("/models/restore-policy", [](const httplib::Request &, httplib::Response & res) {
if (g_keep_loaded_cli) {
json_error(res, 409, "engine was started with --keep-loaded; policy is not restorable");
return;
}
int still = 0;
int freed = store_evict_all(g_store, &still);
bool restored = false;
if (g_keep_loaded && still == 0) {
store_set_policy(g_store, EVICT_STRICT);
g_keep_loaded = false;
restored = true;
fprintf(stderr, "[Server] Eviction policy restored to STRICT (%d module(s) freed)\n", freed);
} else if (still > 0) {
fprintf(stderr, "[Server] restore-policy: %d module(s) still in use — policy left as-is\n", still);
}
char buf[128];
snprintf(buf, sizeof(buf), "{\"freed\":%d,\"resident\":%d,\"restored\":%s}", freed, still,
restored ? "true" : "false");
res.set_content(buf, "application/json");
});
// job system endpoints
svr.Get("/job", [](const httplib::Request & req, httplib::Response & res) {
if (!req.has_param("id")) {
json_error(res, 400, "Missing id parameter");
return;
}
auto job = job_find(req.get_param_value("id"));
if (!job) {
json_error(res, 404, "Job not found");
return;
}
// ?latent=1: return raw post-DiT latent bytes (float32, [T*64])
if (req.has_param("latent") && req.get_param_value("latent") == "1") {
if (job->status.load() != 1 || job->result_latent.empty()) {
json_error(res, 404, "Latent not available");
return;
}
res.set_content(
reinterpret_cast<const char *>(job->result_latent.data()),
job->result_latent.size() * sizeof(float),
"application/octet-stream");
return;
}
// ?result=1: return result body
if (req.has_param("result") && req.get_param_value("result") == "1") {
if (job->status.load() != 1) {
json_error(res, 404, "Result not ready");
return;
}
res.set_content(job->result_body, job->result_mime);
if (!job->result_lrc.empty()) {
res.set_header("X-LRC-Text", job->result_lrc);
}
return;
}
// default: return status JSON. Now includes phase + phase_step/total so
// the wrapper can distinguish "stalled in the ~17 s adapter precompute"
// from "actually failed" without parsing the /logs SSE stream.
char phase_buf[256];
int step = job->phase_step.load(std::memory_order_relaxed);
int total = job->phase_total.load(std::memory_order_relaxed);
snprintf(phase_buf, sizeof(phase_buf),
"{\"status\":\"%s\",\"phase\":\"%s\",\"phase_step\":%d,\"phase_total\":%d}",
job_status_str(job->status.load()), job_phase_str(job->phase.load()), step, total);
res.set_content(phase_buf, "application/json");
});
svr.Post("/job", [](const httplib::Request & req, httplib::Response & res) {
if (!req.has_param("id")) {
json_error(res, 400, "Missing id parameter");
return;
}
auto job = job_find(req.get_param_value("id"));
if (!job) {
json_error(res, 404, "Job not found");
return;
}
// ?cancel=1: cancel the job
if (req.has_param("cancel") && req.get_param_value("cancel") == "1") {
job->cancel.store(true);
fprintf(stderr, "[Server] Cancel requested for job %s\n", job->id.c_str());
res.set_content("{\"status\":\"cancelled\"}", "application/json");
return;
}
json_error(res, 400, "Unknown action");
});
// POST /pp-vae-reencode — synchronous PP-VAE re-encode processing.
// Accepts WAV audio body. Runs PP-VAE encode→decode round-trip with
// RMS gain matching. Returns processed WAV (same sample rate, 16-bit).
// Requires PP-VAE models in registry. Non-fatal: returns 501 if unavailable.
svr.Post("/pp-vae-reencode", [](const httplib::Request & req, httplib::Response & res) {
if (req.body.empty()) {
json_error(res, 400, "Empty body (expected WAV audio)");
return;
}
// Parse blend from query string (0.0 = fully PP-VAE, 1.0 = fully original)
float blend = 0.0f;
if (req.has_param("blend")) {
blend = std::strtof(req.get_param_value("blend").c_str(), nullptr);
if (blend < 0.0f) blend = 0.0f;
if (blend > 1.0f) blend = 1.0f;
}
// Parse backend preference: "onnx" = force ORT/TRT, "gguf" = force GGML, absent = auto
std::string backend = "auto";
if (req.has_param("backend")) {
backend = req.get_param_value("backend");
}
// Resolve PP-VAE model path from registry (prefer F32 > BF16 > F16)
if (g_registry.pp_vae.empty()) {
json_error(res, 501, "No PP-VAE model in registry");
return;
}
const char * pp_vae_path = nullptr;
const char * pref[] = { "F32", "BF16", "F16" };
for (const char * tag : pref) {
for (const auto & e : g_registry.pp_vae) {
if (e.name.find(tag) != std::string::npos) {
pp_vae_path = e.path.c_str();
break;
}
}
if (pp_vae_path) break;
}
if (!pp_vae_path) pp_vae_path = g_registry.pp_vae[0].path.c_str();
// Decode WAV from body → planar stereo [L:T][R:T]
int T_audio = 0;
float * planar = audio_read_48k_buf((const uint8_t *) req.body.data(), req.body.size(), &T_audio);
if (!planar || T_audio <= 0) {
json_error(res, 400, "Failed to decode WAV audio");
return;
}
fprintf(stderr, "[Server] PP-VAE re-encode: %.2fs @ 48kHz, model=%s, blend=%.2f, backend=%s\n",
(float) T_audio / 48000.0f, pp_vae_path, blend, backend.c_str());
// If blend is 1.0 (fully original), skip processing entirely
if (blend >= 1.0f) {
fprintf(stderr, "[Server] PP-VAE: blend=1.0, returning original audio\n");
std::string wav = audio_encode_wav(planar, T_audio, 48000, WAV_S16);
free(planar);
res.set_content(wav, "audio/wav");
return;
}
// Measure input RMS + peak
double in_sum_sq = 0.0;
float in_peak = 0.0f;
int n_total = T_audio * 2;
for (int i = 0; i < n_total; i++) {
float v = planar[i];
in_sum_sq += (double) v * v;
float av = fabsf(v);
if (av > in_peak) in_peak = av;
}
float in_rms = (float) sqrt(in_sum_sq / (double) n_total);
// Resolve PP-VAE ONNX paths for ORT/TRT acceleration.
// Look for pp-vae_encoder.onnx / pp-vae_decoder.onnx in models/onnx/
// Try new subdirectory layout (onnx/pp-vae/) first, fall back to legacy flat layout.
// Skipped entirely when backend=gguf.
std::string pp_dir;
{
std::string p = pp_vae_path;
auto slash = p.find_last_of("/\\");
pp_dir = (slash != std::string::npos) ? p.substr(0, slash) : ".";
}
std::string onnx_dir = pp_dir + "/" + "onnx";
std::string onnx_enc_path, onnx_dec_path;
if (backend != "gguf") {
{
// Try new location first: onnx/pp-vae/pp-vae_encoder.onnx
std::string ep = onnx_dir + "/" + "pp-vae" + "/" + "pp-vae_encoder.onnx";
FILE * f = fopen(ep.c_str(), "rb");
if (!f) {
// Fall back to legacy flat layout
ep = onnx_dir + "/" + "pp-vae_encoder.onnx";
f = fopen(ep.c_str(), "rb");
}
if (f) { fclose(f); onnx_enc_path = ep; }
}
{
// Try new location first: onnx/pp-vae/pp-vae_decoder.onnx
std::string dp = onnx_dir + "/" + "pp-vae" + "/" + "pp-vae_decoder.onnx";
FILE * f = fopen(dp.c_str(), "rb");
if (!f) {
// Fall back to legacy flat layout
dp = onnx_dir + "/" + "pp-vae_decoder.onnx";
f = fopen(dp.c_str(), "rb");
}
if (f) { fclose(f); onnx_dec_path = dp; }
}
if (backend == "onnx" && (onnx_enc_path.empty() || onnx_dec_path.empty())) {
fprintf(stderr, "[Server] PP-VAE backend=onnx but ONNX models not found in %s, falling back to GGML\n",
onnx_dir.c_str());
}
} else {
fprintf(stderr, "[Server] PP-VAE backend=gguf, skipping ONNX discovery\n");
}
// Default VAE tiling params (match scragvae: same Oobleck architecture)
int vae_chunk = 1024;
int vae_overlap = 64;
// Phase 1: Encode (planar → interleaved → VAE encoder → latents)
// Prefers ORT/TRT when pp-vae_encoder.onnx exists, falls back to GGML.
std::vector<float> latents;
int T_latent = 0;
// Convert planar → interleaved for encoder
std::vector<float> interleaved(T_audio * 2);
{
const float * L = planar;
const float * R = planar + T_audio;
for (int i = 0; i < T_audio; i++) {
interleaved[i * 2 + 0] = L[i];
interleaved[i * 2 + 1] = R[i];
}
}
int max_T = (T_audio / 1920) + 64;
latents.resize((size_t) max_T * 64);
if (!onnx_enc_path.empty()) {
// Try ORT encoder
ModelKey enc_ort_key;
enc_ort_key.kind = MODEL_VAE_ENC_ORT;
enc_ort_key.path = onnx_enc_path;
VaeEncOrt * enc_ort = store_require_vae_enc_ort(g_store, enc_ort_key);
if (enc_ort) {
ModelHandle enc_guard(g_store, enc_ort);
fprintf(stderr, "[Server] PP-VAE encoding via ORT/TRT: %s\n", onnx_enc_path.c_str());
T_latent = vae_enc_ort_encode_tiled(enc_ort, interleaved.data(), T_audio,
latents.data(), max_T, vae_chunk, vae_overlap);
} else {
fprintf(stderr, "[Server] PP-VAE ORT encoder load failed, falling back to GGML\n");
}
}
if (T_latent <= 0) {
// Fall back to GGML encoder
ModelKey enc_key;
enc_key.kind = MODEL_VAE_ENC;
enc_key.path = pp_vae_path;
VAEEncoder * enc = store_require_vae_enc(g_store, enc_key);
if (!enc) {
free(planar);
json_error(res, 500, "Failed to load PP-VAE encoder");
return;
}
ModelHandle enc_guard(g_store, enc);
fprintf(stderr, "[Server] PP-VAE encoding via GGML\n");
T_latent = vae_enc_encode_tiled(enc, interleaved.data(), T_audio,
latents.data(), max_T, vae_chunk, vae_overlap);
if (T_latent <= 0) {
free(planar);
json_error(res, 500, "PP-VAE encode failed");
return;
}
}
fprintf(stderr, "[Server] PP-VAE encode: T_latent=%d\n", T_latent);
// Phase 2: Decode (latents → VAE decoder → planar PCM)
// Prefers ORT/TRT when pp-vae_decoder.onnx exists, falls back to GGML.
std::vector<float> decoded;
int T_decoded = 0;
int T_audio_max = T_latent * 1920;
decoded.resize(2 * T_audio_max);
if (!onnx_dec_path.empty()) {
// Try ORT decoder
ModelKey dec_ort_key;
dec_ort_key.kind = MODEL_VAE_DEC_ORT;
dec_ort_key.path = onnx_dec_path;
VaeOrt * dec_ort = store_require_vae_dec_ort(g_store, dec_ort_key);
if (dec_ort) {
ModelHandle dec_guard(g_store, dec_ort);
fprintf(stderr, "[Server] PP-VAE decoding via ORT/TRT: %s\n", onnx_dec_path.c_str());
T_decoded = vae_ort_decode_tiled(dec_ort, latents.data(), T_latent,
decoded.data(), T_audio_max, vae_chunk, vae_overlap);
} else {
fprintf(stderr, "[Server] PP-VAE ORT decoder load failed, falling back to GGML\n");
}
}
if (T_decoded <= 0) {
// Fall back to GGML decoder
ModelKey dec_key;
dec_key.kind = MODEL_VAE_DEC;
dec_key.path = pp_vae_path;
VAEGGML * dec = store_require_vae_dec(g_store, dec_key);
if (!dec) {
free(planar);
json_error(res, 500, "Failed to load PP-VAE decoder");
return;
}
ModelHandle dec_guard(g_store, dec);
fprintf(stderr, "[Server] PP-VAE decoding via GGML\n");
T_decoded = vae_ggml_decode_tiled(dec, latents.data(), T_latent,
decoded.data(), T_audio_max, vae_chunk, vae_overlap, NULL, NULL);
if (T_decoded <= 0) {
free(planar);
json_error(res, 500, "PP-VAE decode failed");
return;
}
}
fprintf(stderr, "[Server] PP-VAE decode: T_decoded=%d\n", T_decoded);
// Phase 3: RMS gain match (scale output to match input RMS, cap at input peak)
double out_sum_sq = 0.0;
float out_peak = 0.0f;
int dec_total = T_decoded * 2;
for (int i = 0; i < dec_total; i++) {
float v = decoded[i];
out_sum_sq += (double) v * v;
float av = fabsf(v);
if (av > out_peak) out_peak = av;
}
float out_rms = (float) sqrt(out_sum_sq / (double) dec_total);
float gain = 1.0f;
if (out_rms > 1e-8f) {
gain = in_rms / out_rms;
if (out_peak * gain > in_peak + 0.01f) {
gain = in_peak / (out_peak + 1e-8f);
}
}
for (int i = 0; i < dec_total; i++) {
decoded[i] *= gain;
}
// Phase 4: Blend original audio into PP-VAE output
// blend=0 → fully PP-VAE, blend=1 → fully original
if (blend > 0.0f) {
int blend_len = std::min(n_total, dec_total);
float wet = 1.0f - blend;
for (int i = 0; i < blend_len; i++) {
decoded[i] = decoded[i] * wet + planar[i] * blend;
}
fprintf(stderr, "[Server] PP-VAE blend: %.0f%% PP-VAE + %.0f%% original\n",
wet * 100.0f, blend * 100.0f);
}
fprintf(stderr, "[Server] PP-VAE done: gain=%.3f (in_rms=%.4f, out_rms=%.4f)\n",
gain, in_rms, out_rms);
free(planar);
// Encode to WAV16 and return
std::string wav = audio_encode_wav(decoded.data(), T_decoded, 48000, WAV_S16);
res.set_content(wav, "audio/wav");
});
// POST /sa3-refine — synchronous SA3 SDEdit refine (instrumental de-fizz).
// Encodes to SAME-L latents, partially re-noises, denoises with the SA3
// DiT, decodes. Numerical reference: tools/onnx-export/e2e_sa3_ort.py.
// Body: WAV or MP3 audio (any sample rate; processed at 44.1k, returned
// at the input rate, 16-bit WAV).
// Query params:
// tokens csv of 256 padded T5Gemma token ids (Node tokenizes;
// bpe.h cannot parse SentencePiece tokenizer.json)
// n_tokens valid (non-pad) token count
// strength init noise level (default 0.3)
// steps sampler steps (default 8)
// sampler "pingpong" (default) | "euler"
// seed RNG seed (default: random)
// rms_match 1 (default) match output RMS to input | 0 raw
// env_match 1 = windowed envelope match to the input (timbre from the
// refine, dynamics from the source); supersedes rms_match
// mix 0..1 wet/dry blend with the SOURCE: 0 = pure source,
// 1 = pure refined (default). Mutually exclusive with the
// band splice below (mix wins when both are sent).
// band_blend 1 = spectral splice: source below the crossover, refined
// above, raised-cosine transition. band_freq = crossover
// center Hz (default 250), band_width = transition width Hz
// (default 200). STFT 8192/hop 2048, Hann, weight-normalized
// overlap-add (linear-phase; no IIR crossover phase seam).
// out_sr output sample rate (default: input rate)
// debug_zero_noise 1 = deterministic validation mode (zero noise)
// backend "onnx" (5 graphs in models/onnx/sa3/) | "gguf" (4 sa3-*.gguf
// in the models root) | "auto" (default: onnx if present,
// else gguf). 501 if the selected backend's models are absent.
// adapters CSV "name:strength,name:strength" — StableStep DoRA adapter
// GGUFs from <models>/sa3-adapters/<name>.gguf, merged into
// the DiT at load. Forces the GGUF backend (ONNX graphs are
// frozen). 400 if a named adapter file is missing.
svr.Post("/sa3-refine", [models_dir](const httplib::Request & req, httplib::Response & res) {
if (req.body.empty()) {
json_error(res, 400, "Empty body (expected WAV audio)");
return;
}
std::string sa3_dir = std::string(models_dir) + "/onnx/sa3";
auto file_exists = [](const std::string & p) {
FILE * f = fopen(p.c_str(), "rb");
if (f) { fclose(f); return true; }
return false;
};
bool have_onnx = file_exists(sa3_dir + "/sa3-dit.onnx");
bool have_gguf = file_exists(std::string(models_dir) + "/sa3-dit-BF16.gguf");
std::string backend = req.has_param("backend") ? req.get_param_value("backend") : "auto";
// StableStep adapters: parse + resolve BEFORE backend selection —
// adapters exist only on the GGUF path, so they force it.
std::vector<std::pair<std::string, float>> adapter_specs; // (path, scale)
std::string adapter_sig; // "path=scale;..." for the ModelKey
if (req.has_param("adapters") && !req.get_param_value("adapters").empty()) {
const std::string csv = req.get_param_value("adapters");
size_t pos = 0;
while (pos < csv.size()) {
size_t comma = csv.find(',', pos);
if (comma == std::string::npos) comma = csv.size();
std::string entry = csv.substr(pos, comma - pos);
pos = comma + 1;
if (entry.empty()) continue;
size_t colon = entry.rfind(':');
std::string name = (colon == std::string::npos) ? entry : entry.substr(0, colon);
float scale = (colon == std::string::npos) ? 1.0f
: strtof(entry.c_str() + colon + 1, nullptr);
// Name sanitation: bare filename stem only (no path traversal)
if (name.empty() || name.find('/') != std::string::npos ||
name.find('\\') != std::string::npos || name.find("..") != std::string::npos) {
json_error(res, 400, "Invalid adapter name");
return;
}
std::string path = std::string(models_dir) + "/sa3-adapters/" + name + ".gguf";
if (!file_exists(path)) {
json_error(res, 400, ("SA3 adapter not found: " + name +
" (expected models/sa3-adapters/" + name + ".gguf)").c_str());
return;
}
adapter_specs.push_back({ path, scale });
if (!adapter_sig.empty()) adapter_sig += ";";
adapter_sig += path + "=" + std::to_string(scale);
}
if (!adapter_specs.empty()) {
if (!have_gguf) {
json_error(res, 501, "SA3 adapters require the GGUF backend (sa3-*.gguf not installed)");
return;
}
if (backend == "onnx") {
json_error(res, 400, "SA3 adapters are GGUF-only — remove backend=onnx or switch to gguf");
return;
}
backend = "gguf";
fprintf(stderr, "[Server] SA3 refine: %zu adapter(s) requested — GGUF backend forced\n",
adapter_specs.size());
}
}
bool use_gguf;
if (backend == "onnx") {
if (!have_onnx) {
json_error(res, 501, "SA3 ONNX models not installed (expected models/onnx/sa3/)");
return;
}
use_gguf = false;
} else if (backend == "gguf") {
if (!have_gguf) {
json_error(res, 501, "SA3 GGUF models not installed (expected sa3-*.gguf in models dir)");
return;
}
use_gguf = true;
} else { // auto
if (have_onnx) use_gguf = false;
else if (have_gguf) use_gguf = true;
else {
json_error(res, 501, "SA3 models not installed (expected models/onnx/sa3/ or sa3-*.gguf)");
return;
}
}
// Params
float strength = 0.3f;
int steps = 8;
bool pingpong = true;
bool zero_noise = false;
bool rms_match = true;
uint64_t seed = (uint64_t)time(nullptr) * 2654435761ull;
if (req.has_param("strength")) {
strength = std::strtof(req.get_param_value("strength").c_str(), nullptr);
if (strength < 0.0f) strength = 0.0f;
if (strength > 1.0f) strength = 1.0f;
}
if (req.has_param("steps")) {
steps = atoi(req.get_param_value("steps").c_str());
if (steps < 1) steps = 1;
if (steps > 64) steps = 64;
}
if (req.has_param("sampler") && req.get_param_value("sampler") == "euler") pingpong = false;
if (req.has_param("seed")) seed = std::strtoull(req.get_param_value("seed").c_str(), nullptr, 10);
if (req.has_param("debug_zero_noise") && req.get_param_value("debug_zero_noise") == "1") zero_noise = true;
if (req.has_param("rms_match") && req.get_param_value("rms_match") == "0") rms_match = false;
// Tokenized prompt (padded to SA3_TOK_LEN)
std::vector<int64_t> ids(SA3_TOK_LEN, 0);
int n_tokens = req.has_param("n_tokens") ? atoi(req.get_param_value("n_tokens").c_str()) : 0;
if (req.has_param("tokens")) {
const std::string & csv = req.get_param_value("tokens");
int idx = 0;
const char * p = csv.c_str();
while (*p && idx < SA3_TOK_LEN) {
ids[idx++] = strtoll(p, nullptr, 10);
const char * comma = strchr(p, ',');
if (!comma) break;
p = comma + 1;
}
}
if (n_tokens <= 0) {
json_error(res, 400, "Missing tokens/n_tokens (tokenized prompt required)");
return;
}
// Decode audio at native rate, resample to 44.1k (planar stereo)
int T_in = 0, sr_in = 0;
float * planar = audio_read_buf((const uint8_t *) req.body.data(), req.body.size(), &T_in, &sr_in);
if (!planar || T_in <= 0) {
json_error(res, 400, "Failed to decode audio");
return;
}
int T44 = 0;
float * p44 = audio_resample(planar, T_in, sr_in, SA3_SR, 2, &T44);
free(planar);
if (!p44 || T44 <= 0) {
free(p44);
json_error(res, 500, "Resample to 44.1k failed");
return;
}
fprintf(stderr, "[Server] SA3 refine: %.2fs @ %dHz, strength=%.2f, steps=%d, sampler=%s, backend=%s\n",
(float) T_in / sr_in, sr_in, strength, steps, pingpong ? "pingpong" : "euler",
use_gguf ? "gguf" : "onnx");
// Input RMS (for gain matching)
double in_sum_sq = 0.0;
for (int i = 0; i < T44 * 2; i++) in_sum_sq += (double) p44[i] * p44[i];
float in_rms = (float) sqrt(in_sum_sq / (double)(T44 * 2));
// Acquire model (selected backend) + run
std::vector<float> out44;
bool ok;
Timer refine_timer;
if (use_gguf) {
ModelKey k{};
k.kind = MODEL_SA3_GGML;
k.path = models_dir; // 4 sa3-*.gguf in the models root
k.adapter_stack = adapter_sig; // "" = stock; else distinct cached model
Sa3GgmlRefine * sa3 = store_require_sa3_ggml(g_store, k);
if (!sa3) {
free(p44);
json_error(res, 500, "SA3 GGML model load failed");
return;
}
ModelHandle guard(g_store, sa3);
refine_timer.reset();
ok = sa3_refine_run_ggml(sa3, p44, T44, ids.data(), n_tokens,
strength, steps, pingpong, seed, zero_noise, out44);
} else {
ModelKey k{};
k.kind = MODEL_SA3_ORT;
k.path = sa3_dir;
Sa3Refine * sa3 = store_require_sa3_ort(g_store, k);
if (!sa3) {
free(p44);
json_error(res, 500, "SA3 model load failed");
return;
}
ModelHandle guard(g_store, sa3);
refine_timer.reset();
ok = sa3_refine_run(sa3, p44, T44, ids.data(), n_tokens,
strength, steps, pingpong, seed, zero_noise, out44);
}
fprintf(stderr, "[Server] SA3 refine compute (%s): %.0f ms\n",
use_gguf ? "gguf" : "onnx", refine_timer.ms());
if (!ok) {
free(p44);
json_error(res, 500, "SA3 refine failed");
return;
}
// env_match=1: windowed envelope match — the refined output's
// short-term RMS is gain-ridden to follow the SOURCE's envelope, so
// the refine changes timbre but not dynamics. Motivation: adapters
// trained on mastered material generate loudness-war density
// ("rectangle" waveforms); a single global RMS gain can fix level but
// not crest factor. ~93 ms windows on a ~46 ms grid, per-sample
// linear gain interpolation, combined-channel gain (stereo balance
// preserved). Replaces the global RMS match when active.
bool env_match = req.has_param("env_match") && req.get_param_value("env_match") == "1";
if (env_match) {
const int64_t hop = 2048, win = 4096; // @44.1k: ~46 ms grid, ~93 ms window
const int64_t n_blocks = (T44 + hop - 1) / hop;
std::vector<float> gains((size_t) n_blocks, 1.0f);
float gmin = 1e9f, gmax = 0.0f;
for (int64_t b = 0; b < n_blocks; b++) {
int64_t s0 = b * hop - (win - hop) / 2;
if (s0 < 0) s0 = 0;
int64_t s1 = s0 + win;
if (s1 > T44) s1 = T44;
double in_sq = 0.0, out_sq = 0.0;
for (int64_t i = s0; i < s1; i++) {
in_sq += (double) p44[i] * p44[i]
+ (double) p44[T44 + i] * p44[T44 + i];
out_sq += (double) out44[(size_t) i] * out44[(size_t) i]
+ (double) out44[(size_t) (T44 + i)] * out44[(size_t) (T44 + i)];
}
int64_t n = (s1 - s0) * 2;
if (n < 1) n = 1;
float bin = (float) sqrt(in_sq / (double) n);
float bout = (float) sqrt(out_sq / (double) n);
float g = bin / (bout + 1e-6f);
if (g > 8.0f) g = 8.0f; // cap: don't amplify refine noise into source-only passages
gains[(size_t) b] = g;
if (g < gmin) gmin = g;
if (g > gmax) gmax = g;
}
for (int64_t i = 0; i < T44; i++) {
int64_t b = i / hop;
float fr = (float) (i - b * hop) / (float) hop;
float g0 = gains[(size_t) b];
float g1 = (b + 1 < n_blocks) ? gains[(size_t) (b + 1)] : g0;
float g = g0 + (g1 - g0) * fr;
for (int ch = 0; ch < 2; ch++) {
float v = out44[(size_t) ch * T44 + i] * g;
out44[(size_t) ch * T44 + i] = v < -1.0f ? -1.0f : (v > 1.0f ? 1.0f : v);
}
}
fprintf(stderr, "[Server] SA3 refine envelope match: %lld blocks, gain %.3f..%.3f\n",
(long long) n_blocks, gmin, gmax);
}
// Source blending — wet/dry mix OR spectral band splice (see docs above).
float mix = -1.0f;
if (req.has_param("mix")) {
mix = std::strtof(req.get_param_value("mix").c_str(), nullptr);
if (mix < 0.0f) mix = 0.0f;
if (mix > 1.0f) mix = 1.0f;
}
if (mix >= 0.0f && mix < 1.0f) {
for (int64_t i = 0; i < (int64_t) 2 * T44; i++) {
float v = p44[i] * (1.0f - mix) + out44[(size_t) i] * mix;
out44[(size_t) i] = v < -1.0f ? -1.0f : (v > 1.0f ? 1.0f : v);
}
fprintf(stderr, "[Server] SA3 refine source mix: %.2f\n", mix);
} else if (mix < 0.0f && req.has_param("band_blend") &&
req.get_param_value("band_blend") == "1") {
float fc = 250.0f, bw = 200.0f;
if (req.has_param("band_freq"))
fc = std::strtof(req.get_param_value("band_freq").c_str(), nullptr);
if (req.has_param("band_width"))
bw = std::strtof(req.get_param_value("band_width").c_str(), nullptr);
if (fc < 40.0f) fc = 40.0f;
if (fc > 16000.0f) fc = 16000.0f;
if (bw < 10.0f) bw = 10.0f;
const float f_lo = fc - bw * 0.5f, f_hi = fc + bw * 0.5f;
// STFT splice: per bin, g = 0 below f_lo (all source), 1 above
// f_hi (all refined), raised cosine between. Hann analysis +
// synthesis, weight-normalized OLA (edge-safe, COLA-free).
const int N = 8192, HOP = 2048, BINS = N / 2 + 1;
std::vector<float> win((size_t) N), gcurve((size_t) BINS);
for (int i = 0; i < N; i++)
win[i] = 0.5f - 0.5f * cosf(2.0f * 3.14159265f * i / N);
for (int b = 0; b < BINS; b++) {
float f = (float) b * SA3_SR / N;
float g;
if (f <= f_lo) g = 0.0f;
else if (f >= f_hi) g = 1.0f;
else g = 0.5f - 0.5f * cosf(3.14159265f * (f - f_lo) / (f_hi - f_lo));
gcurve[b] = g;
}
std::vector<float> frame_s((size_t) N), frame_r((size_t) N), frame_o((size_t) N);
std::vector<sl_detail::Cpx> spec_s((size_t) BINS), spec_r((size_t) BINS);
for (int ch = 0; ch < 2; ch++) {
const float * s = p44 + (size_t) ch * T44;
float * r = out44.data() + (size_t) ch * T44;
std::vector<float> acc((size_t) T44, 0.0f), wsum((size_t) T44, 0.0f);
for (int64_t st = 0; st < T44; st += HOP) {
for (int i = 0; i < N; i++) {
int64_t idx = st + i;
float w = win[i];
frame_s[i] = (idx < T44) ? s[idx] * w : 0.0f;
frame_r[i] = (idx < T44) ? r[idx] * w : 0.0f;
}
sl_detail::rfft(frame_s.data(), spec_s.data(), N);
sl_detail::rfft(frame_r.data(), spec_r.data(), N);
for (int b = 0; b < BINS; b++) {
float g = gcurve[b];
spec_s[b] = sl_detail::Cpx(spec_s[b].re * (1.0f - g) + spec_r[b].re * g,
spec_s[b].im * (1.0f - g) + spec_r[b].im * g);
}
sl_detail::irfft(spec_s.data(), frame_o.data(), N);
for (int i = 0; i < N; i++) {
int64_t idx = st + i;
if (idx >= T44) break;
acc[(size_t) idx] += frame_o[i] * win[i];
wsum[(size_t) idx] += win[i] * win[i];
}
}
for (int64_t i = 0; i < T44; i++) {
float v = acc[(size_t) i] / (wsum[(size_t) i] + 1e-9f);
r[i] = v < -1.0f ? -1.0f : (v > 1.0f ? 1.0f : v);
}
}
fprintf(stderr, "[Server] SA3 refine band splice: source < %.0f Hz, refined > %.0f Hz\n",
f_lo, f_hi);
}
free(p44);
// RMS gain matching (same convention as PP-VAE); superseded by the
// envelope match when env_match=1.
if (!env_match && rms_match && in_rms > 1e-6f) {
double out_sum_sq = 0.0;
for (size_t i = 0; i < out44.size(); i++) out_sum_sq += (double) out44[i] * out44[i];
float out_rms = (float) sqrt(out_sum_sq / (double) out44.size());
if (out_rms > 1e-6f) {
float gain = in_rms / out_rms;
for (size_t i = 0; i < out44.size(); i++) {
float v = out44[i] * gain;
out44[i] = v < -1.0f ? -1.0f : (v > 1.0f ? 1.0f : v);
}
fprintf(stderr, "[Server] SA3 refine gain=%.3f (in_rms=%.4f, out_rms=%.4f)\n",
gain, in_rms, out_rms);
}
}
// Resample to the output rate (default: input rate), encode WAV
int sr_out = sr_in;
if (req.has_param("out_sr")) {
int v = atoi(req.get_param_value("out_sr").c_str());
if (v >= 8000 && v <= 192000) sr_out = v;
}
int T_out = 0;
float * out_native = audio_resample(out44.data(), T44, SA3_SR, sr_out, 2, &T_out);
if (!out_native || T_out <= 0) {
free(out_native);
json_error(res, 500, "Resample to output rate failed");
return;
}
std::string wav = audio_encode_wav(out_native, T_out, sr_out, WAV_S16);
free(out_native);
res.set_content(wav, "audio/wav");
});
// ═══════════════════════════════════════════════════════════════════
// SuperSep: Native stem separation via ONNX Runtime
// ═══════════════════════════════════════════════════════════════════
// Global SuperSep context (lazy-initialized on first request)
static SuperSep * g_supersep = nullptr;
static std::mutex mtx_supersep;
// SuperSep job results (separate from main job pool since stems are large)
struct SuperSepJob {
std::string id;
std::atomic<int> status{0}; // 0=running, 1=done, 2=failed
std::atomic<bool> cancel{false};
float progress{0.0f};
std::string progress_msg;
std::mutex mtx_progress;
SuperSepResult * result{nullptr};
std::string model_dir;
std::string error_msg;
~SuperSepJob() {
if (result) supersep_result_free(result);
}
};
static std::mutex mtx_sep_jobs;
static std::unordered_map<std::string, std::shared_ptr<SuperSepJob>> g_sep_jobs;
// POST /supersep/separate — start async stem separation
// Body: raw WAV or MP3 audio
// Query params: level=0..4 (BASIC/VOCAL_SPLIT/FULL/MAXIMUM/VOCALS_ONLY)
// level=4: BS-RoFormer 2-stem output — Vocals (lead+backing) + Instrumental (mix vocals)
// Returns: {"id": "..."}
svr.Post("/supersep/separate", [models_dir](const httplib::Request & req, httplib::Response & res) {
if (req.body.empty()) {
json_error(res, 400, "Empty body (expected audio)");
return;
}
int level = 0;
if (req.has_param("level")) {
level = atoi(req.get_param_value("level").c_str());
if (level < 0) level = 0;
if (level > SUPERSEP_STABLESTEP) level = SUPERSEP_STABLESTEP;
}
// Decode audio to interleaved stereo 44100 Hz
int T_audio = 0, sr = 0;
float * planar = audio_read_buf((const uint8_t *)req.body.data(), req.body.size(), &T_audio, &sr);
if (!planar || T_audio <= 0) {
json_error(res, 400, "Failed to decode audio");
return;
}
// Resample to 44100 if needed
if (sr != 44100) {
int T_rs = 0;
float * resampled = audio_resample(planar, T_audio, sr, 44100, 2, &T_rs);
free(planar);
if (!resampled) {
json_error(res, 500, "Resample to 44100 failed");
return;
}
planar = resampled;
T_audio = T_rs;
}
// Convert planar to interleaved for SuperSep
float * interleaved = audio_planar_to_interleaved(planar, T_audio);
free(planar);
if (!interleaved) {
json_error(res, 500, "OOM converting to interleaved");
return;
}
// Create job
auto job = std::make_shared<SuperSepJob>();
job->id = job_make_id();
job->model_dir = std::string(models_dir) + "/supersep";
{
std::lock_guard<std::mutex> lock(mtx_sep_jobs);
g_sep_jobs[job->id] = job;
}
int n_frames = T_audio;
SuperSepLevel sep_level = (SuperSepLevel)level;
// Push to work queue (GPU-serialized with DiT/LM jobs)
work_push([job, interleaved, n_frames, sep_level]() {
// Initialize SuperSep if needed
{
std::lock_guard<std::mutex> lock(mtx_supersep);
if (!g_supersep) {
g_supersep = supersep_init(job->model_dir.c_str(), 0);
}
}
if (!g_supersep) {
fprintf(stderr, "[Server] SuperSep init failed\n");
free(interleaved);
job->status.store(2);
return;
}
auto progress_cb = [](int stage, const char *msg, float pct, void *ud) {
auto *j = (SuperSepJob *)ud;
std::lock_guard<std::mutex> lock(j->mtx_progress);
j->progress = pct;
j->progress_msg = msg ? msg : "";
};
auto cancel_cb = [](void *ud) -> bool {
auto *j = (SuperSepJob *)ud;
return j->cancel.load();
};
SuperSepResult *result = supersep_run(
g_supersep, interleaved, n_frames, sep_level,
progress_cb, cancel_cb, (void *)job.get()
);
free(interleaved);
if (result) {
job->result = result;
job->status.store(1);
fprintf(stderr, "[Server] SuperSep job %s done (%d stems)\n",
job->id.c_str(), result->n_stems);
} else {
// Capture the last progress message as the error
{
std::lock_guard<std::mutex> lock(job->mtx_progress);
if (job->error_msg.empty()) {
job->error_msg = job->progress_msg.empty()
? "Unknown error during separation"
: job->progress_msg;
}
}
job->status.store(job->cancel.load() ? 3 : 2);
fprintf(stderr, "[Server] SuperSep job %s failed: %s\n",
job->id.c_str(), job->error_msg.c_str());
}
// Release ONNX sessions to reclaim VRAM immediately
supersep_release_models(g_supersep);
});
fprintf(stderr, "[Server] SuperSep job %s created (level=%d, %.1fs audio)\n",
job->id.c_str(), level, (float)T_audio / 44100.0f);
std::string body = "{\"id\":\"" + job->id + "\"}";
res.set_content(body, "application/json");
});
// GET /supersep/progress?id=... — poll progress
svr.Get("/supersep/progress", [](const httplib::Request & req, httplib::Response & res) {
if (!req.has_param("id")) { json_error(res, 400, "Missing id"); return; }
std::string id = req.get_param_value("id");
std::shared_ptr<SuperSepJob> job;
{
std::lock_guard<std::mutex> lock(mtx_sep_jobs);
auto it = g_sep_jobs.find(id);
if (it == g_sep_jobs.end()) { json_error(res, 404, "Job not found"); return; }
job = it->second;
}
yyjson_mut_doc * doc = yyjson_mut_doc_new(NULL);
yyjson_mut_val * root = yyjson_mut_obj(doc);
yyjson_mut_doc_set_root(doc, root);
int status = job->status.load();
yyjson_mut_obj_add_str(doc, root, "status", job_status_str(status));
{
std::lock_guard<std::mutex> lock(job->mtx_progress);
yyjson_mut_obj_add_real(doc, root, "progress", job->progress);
yyjson_mut_obj_add_str(doc, root, "message", job->progress_msg.c_str());
}
if (status == 1 && job->result) {
yyjson_mut_obj_add_int(doc, root, "n_stems", job->result->n_stems);
}
if (status == 2) {
std::lock_guard<std::mutex> lock2(job->mtx_progress);
if (!job->error_msg.empty()) {
yyjson_mut_obj_add_str(doc, root, "error", job->error_msg.c_str());
}
}
char * json = yyjson_mut_write(doc, 0, NULL);
yyjson_mut_doc_free(doc);
res.set_content(json, "application/json");
free(json);
});
// GET /supersep/result?id=... — get stem list (metadata, not audio)
svr.Get("/supersep/result", [](const httplib::Request & req, httplib::Response & res) {
if (!req.has_param("id")) { json_error(res, 400, "Missing id"); return; }
std::string id = req.get_param_value("id");
std::shared_ptr<SuperSepJob> job;
{
std::lock_guard<std::mutex> lock(mtx_sep_jobs);
auto it = g_sep_jobs.find(id);
if (it == g_sep_jobs.end()) { json_error(res, 404, "Job not found"); return; }
job = it->second;
}
if (job->status.load() != 1 || !job->result) {
json_error(res, 409, "Job not complete");
return;
}
yyjson_mut_doc * doc = yyjson_mut_doc_new(NULL);
yyjson_mut_val * root = yyjson_mut_obj(doc);
yyjson_mut_doc_set_root(doc, root);
yyjson_mut_val * arr = yyjson_mut_arr(doc);
for (int i = 0; i < job->result->n_stems; i++) {
SuperSepStem & s = job->result->stems[i];
yyjson_mut_val * obj = yyjson_mut_obj(doc);
yyjson_mut_obj_add_str(doc, obj, "name", s.name);
yyjson_mut_obj_add_str(doc, obj, "category", s.category);
yyjson_mut_obj_add_str(doc, obj, "stem_type", s.stem_type);
yyjson_mut_obj_add_int(doc, obj, "n_frames", s.n_frames);
yyjson_mut_obj_add_int(doc, obj, "stage", s.stage);
yyjson_mut_obj_add_int(doc, obj, "index", i);
yyjson_mut_obj_add_bool(doc, obj, "hidden", s.hidden);
yyjson_mut_arr_append(arr, obj);
}
yyjson_mut_obj_add_val(doc, root, "stems", arr);
yyjson_mut_obj_add_str(doc, root, "id", id.c_str());
char * json = yyjson_mut_write(doc, 0, NULL);
yyjson_mut_doc_free(doc);
res.set_content(json, "application/json");
free(json);
});
// GET /supersep/serve?id=...&stem=N — download individual stem as WAV
svr.Get("/supersep/serve", [](const httplib::Request & req, httplib::Response & res) {
if (!req.has_param("id") || !req.has_param("stem")) {
json_error(res, 400, "Missing id or stem"); return;
}
std::string id = req.get_param_value("id");
int stem_idx = atoi(req.get_param_value("stem").c_str());
std::shared_ptr<SuperSepJob> job;
{
std::lock_guard<std::mutex> lock(mtx_sep_jobs);
auto it = g_sep_jobs.find(id);
if (it == g_sep_jobs.end()) { json_error(res, 404, "Job not found"); return; }
job = it->second;
}
if (job->status.load() != 1 || !job->result) {
json_error(res, 409, "Job not complete"); return;
}
if (stem_idx < 0 || stem_idx >= job->result->n_stems) {
json_error(res, 400, "Invalid stem index"); return;
}
SuperSepStem & s = job->result->stems[stem_idx];
// Convert interleaved to planar for WAV encoder
float * planar = (float *)malloc(sizeof(float) * s.n_frames * 2);
if (!planar) { json_error(res, 500, "OOM"); return; }
for (int i = 0; i < s.n_frames; i++) {
planar[i] = s.samples[i * 2 + 0];
planar[s.n_frames + i] = s.samples[i * 2 + 1];
}
std::string wav = audio_encode_wav(planar, s.n_frames, 44100, WAV_S16);
free(planar);
res.set_content(wav, "audio/wav");
});
// POST /supersep/release?id=... — drop a job from the pool, freeing its
// stems (each VOCALS_ONLY job pins ~2 full-track float stems in RAM and
// g_sep_jobs has no eviction, so long sessions / batch callers leak
// without this). A still-running worker holds its own shared_ptr, so the
// job is cancelled and memory is reclaimed when the worker finishes.
svr.Post("/supersep/release", [](const httplib::Request & req, httplib::Response & res) {
if (!req.has_param("id")) { json_error(res, 400, "Missing id"); return; }
std::string id = req.get_param_value("id");
std::shared_ptr<SuperSepJob> job;
{
std::lock_guard<std::mutex> lock(mtx_sep_jobs);
auto it = g_sep_jobs.find(id);
if (it == g_sep_jobs.end()) { json_error(res, 404, "Job not found"); return; }
job = it->second;
g_sep_jobs.erase(it);
}
if (job->status.load() == 0) job->cancel.store(true);
res.set_content("{\"released\":true}", "application/json");
});
// POST /supersep/recombine — mix stems with volume/mute, return WAV
// Body: JSON {"id":"...", "stems":[{"index":0,"volume":1.0,"muted":false},...]}
svr.Post("/supersep/recombine", [](const httplib::Request & req, httplib::Response & res) {
yyjson_doc * doc = yyjson_read(req.body.c_str(), req.body.size(), 0);
if (!doc) { json_error(res, 400, "Invalid JSON"); return; }
yyjson_val * root = yyjson_doc_get_root(doc);
yyjson_val * v_id = yyjson_obj_get(root, "id");
if (!v_id) { yyjson_doc_free(doc); json_error(res, 400, "Missing id"); return; }
std::string id = yyjson_get_str(v_id);
std::shared_ptr<SuperSepJob> job;
{
std::lock_guard<std::mutex> lock(mtx_sep_jobs);
auto it = g_sep_jobs.find(id);
if (it == g_sep_jobs.end()) {
yyjson_doc_free(doc);
json_error(res, 404, "Job not found"); return;
}
job = it->second;
}
if (job->status.load() != 1 || !job->result) {
yyjson_doc_free(doc);
json_error(res, 409, "Job not complete"); return;
}
// Parse stem controls
yyjson_val * arr = yyjson_obj_get(root, "stems");
int n = job->result->n_stems;
std::vector<float> volumes(n, 1.0f);
// NB: std::vector<bool> is a packed-bit proxy — no .data().
// Use a real bool array for the C API.
std::unique_ptr<bool[]> muted(new bool[n]());
if (arr && yyjson_is_arr(arr)) {
yyjson_val * item;
size_t idx, max_val;
yyjson_arr_foreach(arr, idx, max_val, item) {
yyjson_val * vi = yyjson_obj_get(item, "index");
if (!vi) continue;
int si = (int)yyjson_get_int(vi);
if (si < 0 || si >= n) continue;
yyjson_val * vv = yyjson_obj_get(item, "volume");
if (vv && yyjson_is_num(vv)) volumes[si] = (float)yyjson_get_num(vv);
yyjson_val * vm = yyjson_obj_get(item, "muted");
if (vm && yyjson_is_bool(vm)) muted[si] = yyjson_get_bool(vm);
}
}
yyjson_doc_free(doc);
// Debug: log the effective mix controls
fprintf(stderr, "[SuperSep] Recombine request: %d stems\n", n);
for (int i = 0; i < n; i++) {
fprintf(stderr, " [%d] %-20s vol=%.2f muted=%d\n",
i, job->result->stems[i].name, volumes[i], (int)muted[i]);
}
int out_frames = 0;
float * mixed = supersep_recombine(
job->result->stems, volumes.data(), muted.get(), n, &out_frames);
if (!mixed || out_frames <= 0) {
json_error(res, 500, "Recombine produced no audio");
return;
}
// Convert interleaved to planar for resampling
float * planar44 = (float *)malloc(sizeof(float) * out_frames * 2);
for (int i = 0; i < out_frames; i++) {
planar44[i] = mixed[i * 2 + 0];
planar44[out_frames + i] = mixed[i * 2 + 1];
}
free(mixed);
// Resample 44100 → 48000 Hz (engine expects 48 kHz)
int out48_frames = 0;
float * planar48 = audio_resample(planar44, out_frames, 44100, 48000, 2, &out48_frames);
free(planar44);
if (!planar48 || out48_frames <= 0) {
json_error(res, 500, "Resample to 48kHz failed");
return;
}
fprintf(stderr, "[SuperSep] Recombined: %d frames @44.1k → %d frames @48k\n",
out_frames, out48_frames);
std::string wav = audio_encode_wav(planar48, out48_frames, 48000, WAV_S16);
free(planar48);
res.set_content(wav, "audio/wav");
});
// POST /spectral-lifter — synchronous Spectral Lifter processing.
// Accepts WAV audio body. SL params are in query string:
// ?denoise_strength=0.3&noise_floor=0.1&hf_mix=0&transient_boost=0&shimmer_reduction=6
// Returns processed WAV audio body (same sample rate, format).
// Runs synchronously (no job queue) — it's pure CPU DSP, typically <1s.
svr.Post("/spectral-lifter", [](const httplib::Request & req, httplib::Response & res) {
if (req.body.empty()) {
json_error(res, 400, "Empty body (expected WAV audio)");
return;
}
// Parse SL params from query string (with defaults)
SpectralLifterParams slp;
spectral_lifter_default(&slp);
if (req.has_param("denoise_strength")) slp.denoise_strength = strtof(req.get_param_value("denoise_strength").c_str(), nullptr);
if (req.has_param("noise_floor")) slp.noise_floor = strtof(req.get_param_value("noise_floor").c_str(), nullptr);
if (req.has_param("hf_mix")) slp.hf_mix = strtof(req.get_param_value("hf_mix").c_str(), nullptr);
if (req.has_param("transient_boost")) slp.transient_boost = strtof(req.get_param_value("transient_boost").c_str(), nullptr);
if (req.has_param("shimmer_reduction")) slp.shimmer_reduction = strtof(req.get_param_value("shimmer_reduction").c_str(), nullptr);
// Decode WAV from body
int T_audio = 0;
float * planar = audio_read_48k_buf((const uint8_t *) req.body.data(), req.body.size(), &T_audio);
if (!planar || T_audio <= 0) {
json_error(res, 400, "Failed to decode WAV audio");
return;
}
fprintf(stderr, "[Server] Spectral Lifter: %.2fs @ 48kHz (denoise=%.2f, floor=%.2f, hf=%.2f, transient=%.2f, shimmer=%.1fdB)\n",
(float) T_audio / 48000.0f, slp.denoise_strength, slp.noise_floor,
slp.hf_mix, slp.transient_boost, slp.shimmer_reduction);
// Process in-place
spectral_lifter_process(planar, T_audio, 48000, &slp);
// Encode back to WAV16
std::string wav = audio_encode_wav(planar, T_audio, 48000, WAV_S16);
free(planar);
res.set_content(wav, "audio/wav");
});
// embedded webui: gzipped single-page app (built by tools/webui/).
// the browser decompresses transparently via Content-Encoding: gzip.
// the .gz is committed to git so cloning + cmake + make gives a working UI.
if (index_html_gz_len > 0) {
svr.Get("/", [](const httplib::Request & req, httplib::Response & res) {
if (req.get_header_value("Accept-Encoding").find("gzip") == std::string::npos) {
res.set_content("Error: gzip is not supported by this browser", "text/plain");
} else {
res.set_header("Content-Encoding", "gzip");
res.set_content(reinterpret_cast<const char *>(index_html_gz), index_html_gz_len,
"text/html; charset=utf-8");
}
});
}
// graceful shutdown on SIGINT/SIGTERM
signal(SIGINT, on_signal);
signal(SIGTERM, on_signal);
// start FIFO worker thread (processes all GPU jobs in order)
std::thread worker(worker_main);
fprintf(stderr, "[Server] acestep.cpp %s\n", ACE_VERSION);
fprintf(stderr, "[Server] Listening on %s:%d\n", host, port);
fprintf(stderr, "[Server] Pipelines:%s%s%s\n", have_lm ? " /lm" : "", have_synth ? " /synth" : "",
have_understand ? " /understand" : "");
fprintf(stderr, "[Server] Models: %zu LM, %zu Text-Enc, %zu DiT, %zu VAE, %zu Adapter\n", g_registry.lm.size(),
g_registry.text_enc.size(), g_registry.dit.size(), g_registry.vae.size(), g_registry.adapters.size());
if (!svr.listen(host, port)) {
fprintf(stderr, "[Server] FATAL: cannot bind %s:%d\n", host, port);
}
// stop worker thread (finishes current job, discards pending)
{
std::lock_guard<std::mutex> lock(mtx_work);
g_work_stop = true;
}
cv_work.notify_one();
worker.join();
// cleanup
fprintf(stderr, "[Server] Shutting down...\n");
store_free(g_store);
fprintf(stderr, "[Server] Done\n");
return 0;
}