// ace-midi.cpp: MuScriptor audio->MIDI transcription (GGML port)
//
// Native port of the MuScriptor transcription model (Kyutai & Mirelo,
// arXiv:2607.08168, code MIT, weights CC BY-NC 4.0). Decoder-only causal
// transformer with mel-spectrogram prefix conditioning; MT3 event vocab.
// Design + validation plan: docs/plans/muscriptor-cpp-port.md.
//
// Phase 1 (this file, current state): weight loading + transformer prefill
// graph + logit-parity selftest against the Python oracle dumps produced by
// tools/ace-midi-validate.py.
//
// ace-midi --model
--validate
// /model.safetensors + config.json ; validation dir with
// prefix.bin / logits_bos.bin / manifest.json
//
// Later phases add: mel frontend, chunked greedy decode with KV cache +
// prelude forcing, note-event decode, MIDI writer, JSONL streaming.
#include
#include
#include
#include
#include
#include
#include
#include "audio-io.h"
#include "backend.h"
#include "ggml.h"
#include "safetensors.h"
// ---------------------------------------------------------------------------
// Config
// ---------------------------------------------------------------------------
struct MidiConfig {
int dim = 768;
int num_heads = 12;
int num_layers = 14;
int card = 1393;
int head_dim() const { return dim / num_heads; }
int ffn_dim() const { return 4 * dim; }
int bos_id() const { return card; } // "initial token" = card
};
static int json_int_field(const char * json, const char * key, int fb) {
char needle[128];
snprintf(needle, sizeof(needle), "\"%s\"", key);
const char * p = strstr(json, needle);
if (!p) return fb;
p = strchr(p + strlen(needle), ':');
if (!p) return fb;
return atoi(p + 1);
}
static bool load_config(MidiConfig * c, const std::string & dir) {
std::string path = dir + "/config.json";
FILE * f = fopen(path.c_str(), "rb");
if (!f) {
fprintf(stderr, "[ace-midi] cannot open %s\n", path.c_str());
return false;
}
std::string j(65536, 0);
size_t n = fread(j.data(), 1, j.size() - 1, f);
fclose(f);
j.resize(n);
c->dim = json_int_field(j.c_str(), "dim", c->dim);
c->num_heads = json_int_field(j.c_str(), "num_heads", c->num_heads);
c->num_layers = json_int_field(j.c_str(), "num_layers", c->num_layers);
c->card = json_int_field(j.c_str(), "card", c->card);
fprintf(stderr, "[ace-midi] config: dim=%d heads=%d layers=%d card=%d\n",
c->dim, c->num_heads, c->num_layers, c->card);
return true;
}
// ---------------------------------------------------------------------------
// Model weights
// ---------------------------------------------------------------------------
struct MidiLayer {
ggml_tensor * norm1_w, * norm1_b;
ggml_tensor * in_proj; // [dim, 3*dim] (ggml: ne0=in)
ggml_tensor * out_proj; // [dim, dim]
ggml_tensor * norm2_w, * norm2_b;
ggml_tensor * ffn1; // [dim, 4*dim]
ggml_tensor * ffn2; // [4*dim, dim]
};
// Constants fixed by the upstream model (transcription_model.py)
#define MIDI_SAMPLE_RATE 16000
#define MIDI_CHUNK_SAMPLES 80000 // 5 s
#define MIDI_N_FFT 2048
#define MIDI_HOP 160 // 100 Hz frame rate
#define MIDI_N_MELS 512
#define MIDI_MEL_FRAMES 501 // 1 + 80000/160 (center=True)
#define MIDI_MAX_GEN 2000 // max tokens per chunk
#define MIDI_EOS_ID 1
struct MidiModel {
MidiConfig cfg;
ggml_context * wctx = nullptr;
ggml_backend_buffer_t wbuf = nullptr;
ggml_tensor * emb; // [dim, card+1]
std::vector layers;
ggml_tensor * out_norm_w, * out_norm_b;
ggml_tensor * head; // [dim, card]
ggml_backend_t backend, cpu_backend;
ggml_backend_sched_t sched;
// KV cache (batch=1): per layer K [D, max_seq, H], V [max_seq, D, H].
// F32 on CPU (byte-exact oracle parity); F16 on GPU — the f32->f16 cpy
// kernels are the llama.cpp-exercised path (the f32->f32 non-contiguous
// cpy miswrites on CUDA), and f16 KV matches upstream's GPU autocast.
ggml_context * kv_ctx = nullptr;
ggml_backend_buffer_t kv_buf = nullptr;
std::vector kv_k, kv_v;
ggml_type kv_type = GGML_TYPE_F32;
int max_seq = 0;
// Host-side copies for CPU input assembly / mel frontend
std::vector emb_host; // [card+1, dim] (row-major, incl. BOS row)
std::vector mel_window; // [2048]
std::vector mel_fb; // [1025, 512] row-major
std::vector mel_proj_w; // [dim, 512] row-major (torch [out,in])
std::vector mel_proj_b; // [dim]
std::vector ds_null_emb; // dataset_name embed row 1 (None cond)
std::vector ig_null_emb; // instrument_group embed row 1 (None cond)
};
// Create a ggml tensor mirroring a safetensors entry (torch [out,in] row-major
// -> ggml [in, out]) and upload its data, converting BF16/F16 -> F32.
static ggml_tensor * load_tensor(MidiModel * m, ggml_context * ctx, const STFile & st,
const std::string & name, int64_t ne0, int64_t ne1) {
(void) m;
const STEntry * e = st_find(st, name.c_str());
if (!e) {
fprintf(stderr, "[ace-midi] FATAL: missing tensor %s\n", name.c_str());
exit(1);
}
// torch shape is [out, in] row-major -> ggml [ne0=in, ne1=out], same memory
ggml_tensor * t = ne1 > 0
? ggml_new_tensor_2d(ctx, GGML_TYPE_F32, ne0, ne1)
: ggml_new_tensor_1d(ctx, GGML_TYPE_F32, ne0);
ggml_set_name(t, name.c_str());
return t;
}
static void upload_tensor(MidiModel * m, const STFile & st, ggml_tensor * t) {
const STEntry * e = st_find(st, t->name);
size_t n = ggml_nelements(t);
// verify element count matches
int64_t st_n = 1;
for (int i = 0; i < e->n_dims; i++) st_n *= e->shape[i];
if ((int64_t) n != st_n) {
fprintf(stderr, "[ace-midi] FATAL: %s shape mismatch (st=%lld ggml=%zu)\n",
t->name, (long long) st_n, n);
exit(1);
}
const void * src = st_data(st, *e);
if (e->dtype == "F32") {
ggml_backend_tensor_set(t, src, 0, n * 4);
} else if (e->dtype == "BF16") {
std::vector tmp(n);
const uint16_t * s = (const uint16_t *) src;
for (size_t i = 0; i < n; i++) {
uint32_t bits = (uint32_t) s[i] << 16;
memcpy(&tmp[i], &bits, 4);
}
ggml_backend_tensor_set(t, tmp.data(), 0, n * 4);
} else if (e->dtype == "F16") {
std::vector tmp(n);
const ggml_fp16_t * s = (const ggml_fp16_t *) src;
for (size_t i = 0; i < n; i++) tmp[i] = ggml_fp16_to_fp32(s[i]);
ggml_backend_tensor_set(t, tmp.data(), 0, n * 4);
} else {
fprintf(stderr, "[ace-midi] FATAL: %s unsupported dtype %s\n", t->name, e->dtype.c_str());
exit(1);
}
}
// Published checkpoints use the legacy multi-codebook key layout for the
// embedding and head (emb.0.* / linears.0.*) — same remap as the Python
// loader's _remap_single_codebook_keys.
static std::string resolve_key(const STFile & st, const std::string & canonical, const std::string & legacy) {
if (st_find(st, canonical.c_str())) return canonical;
if (st_find(st, legacy.c_str())) return legacy;
return canonical; // load_tensor will report it missing
}
static bool load_model(MidiModel * m, const std::string & dir) {
if (!load_config(&m->cfg, dir)) return false;
const MidiConfig & c = m->cfg;
STFile st;
std::string wpath = dir + "/model.safetensors";
if (!st_open(&st, wpath.c_str())) return false;
BackendPair bp = backend_init("MIDI");
m->backend = bp.backend;
m->cpu_backend = bp.cpu_backend;
m->sched = backend_sched_new(bp, 8192);
int n_tensors = 3 /*emb, out_norm w/b*/ + 1 /*head*/ + c.num_layers * 8;
ggml_init_params ip = { (size_t) n_tensors * ggml_tensor_overhead() + 4096, NULL, true };
m->wctx = ggml_init(ip);
const std::string emb_key = resolve_key(st, "emb.weight", "emb.0.weight");
const std::string head_key = resolve_key(st, "linear.weight", "linears.0.weight");
m->emb = load_tensor(m, m->wctx, st, emb_key, c.dim, c.card + 1);
m->out_norm_w = load_tensor(m, m->wctx, st, "out_norm.weight", c.dim, 0);
m->out_norm_b = load_tensor(m, m->wctx, st, "out_norm.bias", c.dim, 0);
m->head = load_tensor(m, m->wctx, st, head_key, c.dim, c.card);
m->layers.resize(c.num_layers);
for (int l = 0; l < c.num_layers; l++) {
char base[96];
snprintf(base, sizeof(base), "transformer.layers.%d.", l);
MidiLayer & L = m->layers[l];
L.norm1_w = load_tensor(m, m->wctx, st, std::string(base) + "norm1.weight", c.dim, 0);
L.norm1_b = load_tensor(m, m->wctx, st, std::string(base) + "norm1.bias", c.dim, 0);
L.in_proj = load_tensor(m, m->wctx, st, std::string(base) + "self_attn.in_proj_weight", c.dim, 3 * c.dim);
L.out_proj = load_tensor(m, m->wctx, st, std::string(base) + "self_attn.out_proj.weight", c.dim, c.dim);
L.norm2_w = load_tensor(m, m->wctx, st, std::string(base) + "norm2.weight", c.dim, 0);
L.norm2_b = load_tensor(m, m->wctx, st, std::string(base) + "norm2.bias", c.dim, 0);
L.ffn1 = load_tensor(m, m->wctx, st, std::string(base) + "linear1.weight", c.dim, c.ffn_dim());
L.ffn2 = load_tensor(m, m->wctx, st, std::string(base) + "linear2.weight", c.ffn_dim(), c.dim);
}
m->wbuf = ggml_backend_alloc_ctx_tensors(m->wctx, m->backend);
if (!m->wbuf) {
fprintf(stderr, "[ace-midi] FATAL: weight buffer alloc failed\n");
return false;
}
for (ggml_tensor * t = ggml_get_first_tensor(m->wctx); t; t = ggml_get_next_tensor(m->wctx, t)) {
upload_tensor(m, st, t);
}
// Host-side copies: token embeddings (input assembly), mel frontend
// weights, and the null class-conditioner rows. Conditioner tokenize
// maps None -> -1, +1 in tokenize, +1 again in forward => row 1.
auto read_host = [&](const char * name, std::vector & out) {
const STEntry * e = st_find(st, name);
if (!e) {
fprintf(stderr, "[ace-midi] FATAL: missing tensor %s\n", name);
exit(1);
}
int64_t n = 1;
for (int i = 0; i < e->n_dims; i++) n *= e->shape[i];
out.resize((size_t) n);
const void * src = st_data(st, *e);
if (e->dtype == "F32") {
memcpy(out.data(), src, (size_t) n * 4);
} else {
const uint16_t * s = (const uint16_t *) src;
for (int64_t i = 0; i < n; i++) {
if (e->dtype == "BF16") {
uint32_t bits = (uint32_t) s[i] << 16;
memcpy(&out[i], &bits, 4);
} else {
out[i] = ggml_fp16_to_fp32((ggml_fp16_t) s[i]);
}
}
}
};
read_host(emb_key.c_str(), m->emb_host);
read_host("condition_provider.conditioners.self_wav.mel_spec_transform.spectrogram.window", m->mel_window);
read_host("condition_provider.conditioners.self_wav.mel_spec_transform.mel_scale.fb", m->mel_fb);
read_host("condition_provider.conditioners.self_wav.output_proj.weight", m->mel_proj_w);
read_host("condition_provider.conditioners.self_wav.output_proj.bias", m->mel_proj_b);
{
std::vector tmp;
read_host("condition_provider.conditioners.dataset_name.embed.weight", tmp);
m->ds_null_emb.assign(tmp.begin() + c.dim, tmp.begin() + 2 * c.dim); // row 1
read_host("condition_provider.conditioners.instrument_group.embed.weight", tmp);
m->ig_null_emb.assign(tmp.begin() + c.dim, tmp.begin() + 2 * c.dim); // row 1
}
// KV cache: prefix (mel 501 + 2 class conds) + BOS + tie prompt + max gen
m->max_seq = MIDI_MEL_FRAMES + 2 + 1 + 300 + MIDI_MAX_GEN;
{
const int D = c.head_dim(), H = c.num_heads;
m->kv_type = (bp.backend != bp.cpu_backend) ? GGML_TYPE_F16 : GGML_TYPE_F32;
ggml_init_params kp = { (size_t) c.num_layers * 2 * ggml_tensor_overhead() + 4096, NULL, true };
m->kv_ctx = ggml_init(kp);
m->kv_k.resize(c.num_layers);
m->kv_v.resize(c.num_layers);
for (int l = 0; l < c.num_layers; l++) {
m->kv_k[l] = ggml_new_tensor_3d(m->kv_ctx, m->kv_type, D, m->max_seq, H);
m->kv_v[l] = ggml_new_tensor_3d(m->kv_ctx, m->kv_type, m->max_seq, D, H);
char nm[32];
snprintf(nm, sizeof(nm), "kv_k_%d", l);
ggml_set_name(m->kv_k[l], nm);
snprintf(nm, sizeof(nm), "kv_v_%d", l);
ggml_set_name(m->kv_v[l], nm);
}
m->kv_buf = ggml_backend_alloc_ctx_tensors(m->kv_ctx, m->backend);
if (!m->kv_buf) {
fprintf(stderr, "[ace-midi] FATAL: KV cache alloc failed\n");
return false;
}
}
fprintf(stderr, "[ace-midi] loaded %d layers (%.1f MB weights, %.1f MB KV cache)\n",
c.num_layers, (double) ggml_backend_buffer_get_size(m->wbuf) / 1e6,
(double) ggml_backend_buffer_get_size(m->kv_buf) / 1e6);
st_close(&st);
return true;
}
// ---------------------------------------------------------------------------
// Sinusoidal positions (transformer.py create_sin_embedding: cat([cos, sin]),
// exponent i/(half_dim - 1), max_period 10000)
// ---------------------------------------------------------------------------
static void add_sin_pos(float * x, int T, int dim, int pos0) {
int half = dim / 2;
for (int t = 0; t < T; t++) {
double pos = (double) (pos0 + t);
for (int i = 0; i < half; i++) {
double phase = pos / pow(10000.0, (double) i / (double) (half - 1));
x[(size_t) t * dim + i] += (float) cos(phase);
x[(size_t) t * dim + half + i] += (float) sin(phase);
}
}
}
// ---------------------------------------------------------------------------
// Prefill forward: input embeddings [dim, T] -> logits [card, T]
// ---------------------------------------------------------------------------
static ggml_tensor * build_layer_norm(ggml_context * ctx, ggml_tensor * x,
ggml_tensor * w, ggml_tensor * b) {
x = ggml_norm(ctx, x, 1e-5f);
x = ggml_mul(ctx, x, w);
return ggml_add(ctx, x, b);
}
// Debug probe (env ACE_MIDI_PROBE=): on the first T==1 step,
// dump layer-0 intermediates so backends can be diffed op by op.
struct ProbeSlot { const char * name; ggml_tensor * t; };
static std::vector g_probe_slots;
static bool g_probe_armed = false;
// Forward T tokens at cache position n_past; writes K/V into the cache and
// reads back the LAST position's logits. n_past=0 with T>1 is the prefill;
// T=1 with n_past>0 is a decode step. Caller advances n_past by T afterwards.
static void forward_tokens(MidiModel * m, const float * input, int T, int n_past, float * logits_last) {
const MidiConfig & c = m->cfg;
const int H = c.num_heads, D = c.head_dim();
const int n_kv = n_past + T;
ggml_init_params ip = { ggml_tensor_overhead() * 8192 + ggml_graph_overhead_custom(8192, false), NULL, true };
ggml_context * ctx = ggml_init(ip);
ggml_tensor * inp = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, c.dim, T);
ggml_set_name(inp, "inp");
ggml_set_input(inp);
ggml_cgraph * gf = ggml_new_graph_custom(ctx, 8192, false);
static int step_calls = 0;
const bool probe = getenv("ACE_MIDI_PROBE") && T == 1 && step_calls == 0;
if (T == 1) step_calls++;
g_probe_slots.clear();
auto probe_add = [&](const char * name, ggml_tensor * t) {
if (probe) { ggml_set_output(t); g_probe_slots.push_back({ name, t }); }
};
ggml_tensor * x = inp;
for (int l = 0; l < c.num_layers; l++) {
MidiLayer & L = m->layers[l];
// --- causal self-attention with KV cache ---
ggml_tensor * h = build_layer_norm(ctx, x, L.norm1_w, L.norm1_b);
ggml_tensor * qkv = ggml_mul_mat(ctx, L.in_proj, h); // [3*dim, T]
ggml_tensor * q = ggml_view_2d(ctx, qkv, c.dim, T, qkv->nb[1], 0);
ggml_tensor * k = ggml_view_2d(ctx, qkv, c.dim, T, qkv->nb[1], (size_t) c.dim * 4);
ggml_tensor * v = ggml_view_2d(ctx, qkv, c.dim, T, qkv->nb[1], (size_t) 2 * c.dim * 4);
// packed layout per token is [h, d] (rearrange "(p h d)")
ggml_tensor * k3 = ggml_reshape_3d(ctx, ggml_cont(ctx, k), D, H, T);
ggml_tensor * v3 = ggml_reshape_3d(ctx, ggml_cont(ctx, v), D, H, T);
// append current K rows: cache K layout [D, max_seq, H], slice dim1 [n_past, n_past+T)
// append current K rows: cache K layout [D, max_seq, H], slice dim1
ggml_tensor * kc = m->kv_k[l];
ggml_tensor * k_dst = ggml_view_3d(ctx, kc, D, T, H, kc->nb[1], kc->nb[2],
(size_t) n_past * kc->nb[1]);
ggml_build_forward_expand(gf, ggml_cpy(ctx, ggml_permute(ctx, k3, 0, 2, 1, 3), k_dst));
// append current V rows into the transposed cache [max_seq, D, H].
// Write via the llama.cpp idiom — 2-D transposed src into a 2-D
// strided view — the one scatter pattern the CUDA cpy kernel is
// known-good for. The earlier 3-D single-element-per-row view write
// silently corrupted rows on CUDA (heads >= 1 got stale VRAM), which
// wrecked every decode step; CPU was unaffected. Probe-verified.
ggml_tensor * vc = m->kv_v[l];
ggml_tensor * v_dst = ggml_view_2d(ctx, vc, T, D * H, vc->nb[1],
(size_t) n_past * vc->nb[0]);
ggml_tensor * v2 = ggml_cont(ctx, v); // [dim, T]
ggml_build_forward_expand(gf, ggml_cpy(ctx, ggml_transpose(ctx, v2), v_dst));
ggml_tensor * Q = ggml_permute(ctx, ggml_reshape_3d(ctx, ggml_cont(ctx, q), D, H, T), 0, 2, 1, 3); // [D, T, H]
ggml_tensor * K = ggml_view_3d(ctx, kc, D, n_kv, H, kc->nb[1], kc->nb[2], 0); // [D, n_kv, H]
ggml_tensor * V = ggml_view_3d(ctx, vc, n_kv, D, H, vc->nb[1], vc->nb[2], 0); // [n_kv, D, H]
ggml_tensor * kq = ggml_mul_mat(ctx, K, Q); // [n_kv, T, H]
kq = ggml_scale(ctx, kq, 1.0f / sqrtf((float) D));
// Causal mask (bottom-right aligned). For T=1 a single query row
// attends the whole cache, so the mask is a mathematical no-op —
// skip it: the legacy diag_mask_inf CUDA kernel miscomputes with
// n_past > 0, which silently corrupted every decode step on GPU.
if (T > 1) kq = ggml_diag_mask_inf(ctx, kq, n_past);
kq = ggml_soft_max(ctx, kq);
ggml_tensor * kqv = ggml_mul_mat(ctx, V, kq); // [D, T, H]
ggml_tensor * att = ggml_cont(ctx, ggml_permute(ctx, kqv, 0, 2, 1, 3)); // [D, H, T]
att = ggml_reshape_2d(ctx, att, c.dim, T);
att = ggml_mul_mat(ctx, L.out_proj, att);
if (l == 0) {
probe_add("h", h);
probe_add("qkv", qkv);
probe_add("K", K);
probe_add("V", V);
probe_add("kqsoft", kq);
probe_add("att", att);
}
x = ggml_add(ctx, x, att);
// --- FFN (exact GELU, matching torch F.gelu default) ---
ggml_tensor * f = build_layer_norm(ctx, x, L.norm2_w, L.norm2_b);
f = ggml_mul_mat(ctx, L.ffn1, f);
f = ggml_gelu_erf(ctx, f);
f = ggml_mul_mat(ctx, L.ffn2, f);
x = ggml_add(ctx, x, f);
}
x = build_layer_norm(ctx, x, m->out_norm_w, m->out_norm_b);
ggml_tensor * logits = ggml_mul_mat(ctx, m->head, x); // [card, T]
ggml_set_name(logits, "logits");
ggml_set_output(logits);
ggml_build_forward_expand(gf, logits);
ggml_backend_sched_reset(m->sched);
if (!ggml_backend_sched_alloc_graph(m->sched, gf)) {
fprintf(stderr, "[ace-midi] FATAL: graph alloc failed\n");
exit(1);
}
ggml_backend_tensor_set(inp, input, 0, (size_t) c.dim * T * 4);
if (ggml_backend_sched_graph_compute(m->sched, gf) != GGML_STATUS_SUCCESS) {
fprintf(stderr, "[ace-midi] FATAL: graph compute failed\n");
exit(1);
}
ggml_backend_tensor_get(logits, logits_last, (size_t) (T - 1) * logits->nb[1], (size_t) c.card * 4);
if (probe) {
const char * prefix_env = getenv("ACE_MIDI_PROBE");
for (auto & s : g_probe_slots) {
char path[512];
snprintf(path, sizeof(path), "%s_%s.bin", prefix_env, s.name);
std::vector buf(ggml_nbytes(s.t));
ggml_backend_tensor_get(s.t, buf.data(), 0, buf.size());
FILE * f = fopen(path, "wb");
if (f) { fwrite(buf.data(), 1, buf.size(), f); fclose(f); }
fprintf(stderr, "[probe] %s: ne=[%lld,%lld,%lld] -> %s\n", s.name,
(long long) s.t->ne[0], (long long) s.t->ne[1], (long long) s.t->ne[2], path);
}
}
ggml_free(ctx);
}
// ---------------------------------------------------------------------------
// Mel frontend (conditioners.py MelSpectrogramConditioner, torchaudio-equiv):
// magnitude STFT (2048/160, periodic hann from ckpt, center reflect pad) ->
// HTK mel fb from ckpt -> log(+1e-6) -> output_proj -> zero masked frames.
// ---------------------------------------------------------------------------
static void fft_radix2(float * re, float * im, int n) {
// exact twiddle table (per stage), computed once — the naive multiplicative
// twiddle recurrence drifts enough to visibly perturb the log-mel
static std::vector tw_re, tw_im;
static int tw_n = 0;
if (tw_n != n) {
tw_re.assign((size_t) n, 0.0f);
tw_im.assign((size_t) n, 0.0f);
for (int len = 2, base = 0; len <= n; len <<= 1, base += len >> 2) {
for (int j = 0; j < len / 2; j++) {
double ang = -2.0 * 3.14159265358979323846 * j / len;
tw_re[(size_t) base + j] = (float) cos(ang);
tw_im[(size_t) base + j] = (float) sin(ang);
}
}
tw_n = n;
}
// bit-reversal permutation
for (int i = 1, j = 0; i < n; i++) {
int bit = n >> 1;
for (; j & bit; bit >>= 1) j ^= bit;
j |= bit;
if (i < j) {
float t;
t = re[i]; re[i] = re[j]; re[j] = t;
t = im[i]; im[i] = im[j]; im[j] = t;
}
}
for (int len = 2, base = 0; len <= n; len <<= 1, base += len >> 2) {
for (int i = 0; i < n; i += len) {
for (int j = 0; j < len / 2; j++) {
int a = i + j, b = i + j + len / 2;
float cr = tw_re[(size_t) base + j], ci = tw_im[(size_t) base + j];
float xr = re[b] * cr - im[b] * ci;
float xi = re[b] * ci + im[b] * cr;
re[b] = re[a] - xr; im[b] = im[a] - xi;
re[a] += xr; im[a] += xi;
}
}
}
}
// Compute the full conditioning prefix for one 5 s chunk:
// [mel 501 | dataset_name(None) | instrument_group(None)] rows of dim floats.
// n_samples is the unpadded chunk length (masks trailing mel frames).
static std::vector compute_prefix(MidiModel * m, const float * wav, int n_samples) {
const MidiConfig & c = m->cfg;
const int pad = MIDI_N_FFT / 2;
const int n = MIDI_CHUNK_SAMPLES;
// reflect-padded chunk (zero-pad the tail to 5 s first, like F.pad)
std::vector padded(n + 2 * pad, 0.0f);
auto sample_at = [&](int i) -> float {
// reflect at both edges of the zero-padded 80000-sample chunk
if (i < 0) i = -i;
if (i >= n) i = 2 * n - 2 - i;
return (i >= 0 && i < n_samples) ? wav[i] : 0.0f;
};
for (int i = 0; i < n + 2 * pad; i++) padded[i] = sample_at(i - pad);
const int T = MIDI_MEL_FRAMES;
std::vector prefix((size_t) (T + 2) * c.dim, 0.0f);
// frames masked at index >= n_samples/160.0 (length_to_mask semantics)
const double frame_limit = (double) n_samples / (double) MIDI_HOP;
std::vector re(MIDI_N_FFT), im(MIDI_N_FFT), logmel(MIDI_N_MELS);
std::vector melacc(MIDI_N_MELS);
for (int f = 0; f < T; f++) {
if ((double) f >= frame_limit) continue; // masked -> zero row
const float * frame = padded.data() + (size_t) f * MIDI_HOP;
for (int i = 0; i < MIDI_N_FFT; i++) {
re[i] = frame[i] * m->mel_window[i];
im[i] = 0.0f;
}
fft_radix2(re.data(), im.data(), MIDI_N_FFT);
// magnitude (power=1.0) -> mel -> log (double accumulation)
const int n_freq = MIDI_N_FFT / 2 + 1;
for (int mm = 0; mm < MIDI_N_MELS; mm++) melacc[mm] = 0.0;
for (int k = 0; k < n_freq; k++) {
double mag = sqrt((double) re[k] * re[k] + (double) im[k] * im[k]);
if (mag == 0.0) continue;
const float * fbrow = m->mel_fb.data() + (size_t) k * MIDI_N_MELS;
for (int mm = 0; mm < MIDI_N_MELS; mm++) melacc[mm] += mag * fbrow[mm];
}
for (int mm = 0; mm < MIDI_N_MELS; mm++) logmel[mm] = logf((float) melacc[mm] + 1e-6f);
// output_proj: [dim, 512] @ logmel + bias
float * out = prefix.data() + (size_t) f * c.dim;
for (int o = 0; o < c.dim; o++) {
const float * wrow = m->mel_proj_w.data() + (size_t) o * MIDI_N_MELS;
double acc = m->mel_proj_b[o];
for (int i = 0; i < MIDI_N_MELS; i++) acc += (double) wrow[i] * logmel[i];
out[o] = (float) acc;
}
}
// class conditioner rows (always the None/null class at inference)
memcpy(prefix.data() + (size_t) T * c.dim, m->ds_null_emb.data(), (size_t) c.dim * 4);
memcpy(prefix.data() + (size_t) (T + 1) * c.dim, m->ig_null_emb.data(), (size_t) c.dim * 4);
return prefix;
}
// ---------------------------------------------------------------------------
// Greedy chunk decode: prefill [prefix | BOS | prompt] then argmax steps.
// Emits every accepted token (prompt tokens included, EOS excluded) via cb.
// ---------------------------------------------------------------------------
static void greedy_argmax_range(const float * logits, int n_valid, int * out) {
int best = 0;
for (int i = 1; i < n_valid; i++) {
if (logits[i] > logits[best]) best = i;
}
*out = best;
}
template
static void decode_chunk(MidiModel * m, const std::vector & prefix,
const std::vector & prompt, int max_gen, TokenCb cb) {
const MidiConfig & c = m->cfg;
const int n_valid = c.card < 1393 ? c.card : 1393; // logits[1393:] masked upstream
const int T_prefix = (int) (prefix.size() / c.dim);
// prefill input: prefix + BOS + prompt tokens, sinusoidal positions from 0
int T0 = T_prefix + 1 + (int) prompt.size();
std::vector input((size_t) T0 * c.dim);
memcpy(input.data(), prefix.data(), prefix.size() * 4);
memcpy(input.data() + prefix.size(), m->emb_host.data() + (size_t) c.bos_id() * c.dim, (size_t) c.dim * 4);
for (size_t i = 0; i < prompt.size(); i++) {
memcpy(input.data() + prefix.size() + (i + 1) * c.dim,
m->emb_host.data() + (size_t) prompt[i] * c.dim, (size_t) c.dim * 4);
cb(prompt[i]); // teacher-forced tokens flow through the stream
}
add_sin_pos(input.data(), T0, c.dim, 0);
std::vector logits(c.card);
forward_tokens(m, input.data(), T0, 0, logits.data());
int n_past = T0;
int tok;
greedy_argmax_range(logits.data(), n_valid, &tok);
std::vector step(c.dim);
for (int i = (int) prompt.size(); i < max_gen; i++) {
if (tok == MIDI_EOS_ID) return;
cb(tok);
if (n_past + 1 > m->max_seq) {
fprintf(stderr, "[ace-midi] WARNING: KV cache full at %d tokens\n", n_past);
return;
}
memcpy(step.data(), m->emb_host.data() + (size_t) tok * c.dim, (size_t) c.dim * 4);
add_sin_pos(step.data(), 1, c.dim, n_past);
forward_tokens(m, step.data(), 1, n_past, logits.data());
n_past++;
greedy_argmax_range(logits.data(), n_valid, &tok);
}
}
// ---------------------------------------------------------------------------
// MT3 event vocabulary (tokenizer/notes.py build_event_vocab, max_shift 1001):
// 0-2 PAD/EOS/UNK | 3-1003 shift | 1004-1131 pitch | 1132-1133 velocity |
// 1134 tie | 1135-1264 program(0-129) | 1265-1392 drum
// ---------------------------------------------------------------------------
enum EvType { EV_SPECIAL, EV_SHIFT, EV_PITCH, EV_VELOCITY, EV_TIE, EV_PROGRAM, EV_DRUM };
struct Ev { EvType type; int value; };
static Ev vocab_decode(int id) {
if (id < 3) return { EV_SPECIAL, id };
if (id < 1004) return { EV_SHIFT, id - 3 };
if (id < 1132) return { EV_PITCH, id - 1004 };
if (id < 1134) return { EV_VELOCITY, id - 1132 };
if (id == 1134) return { EV_TIE, 0 };
if (id < 1265) return { EV_PROGRAM, id - 1135 };
if (id < 1393) return { EV_DRUM, id - 1265 };
return { EV_SPECIAL, 2 };
}
static int tok_program(int program) { return 1135 + program; }
static int tok_pitch(int pitch) { return 1004 + pitch; }
#define TOK_TIE 1134
#define DRUM_PROGRAM 128
#define MIN_NOTE_DUR 0.01
#define FRAME_RATE 100
// MT3_FULL_PLUS named groups: representative (first) program -> group name.
// The model always emits the representative program of a group (mt3.py).
static const char * instrument_for_program(int program) {
switch (program) {
case 0: return "acoustic_piano";
case 2: return "electric_piano";
case 8: return "chromatic_percussion";
case 16: return "organ";
case 24: return "acoustic_guitar";
case 26: return "clean_electric_guitar";
case 29: return "distorted_electric_guitar";
case 32: return "acoustic_bass";
case 33: return "electric_bass";
case 40: return "violin";
case 41: return "viola";
case 42: return "cello";
case 43: return "contrabass";
case 46: return "orchestral_harp";
case 47: return "timpani";
case 48: return "string_ensemble";
case 50: return "synth_strings";
case 52: return "voice";
case 55: return "orchestra_hit";
case 56: return "trumpet";
case 57: return "trombone";
case 58: return "tuba";
case 60: return "french_horn";
case 61: return "brass_section";
case 64: return "soprano_and_alto_sax";
case 66: return "tenor_sax";
case 67: return "baritone_sax";
case 68: return "oboe";
case 69: return "english_horn";
case 70: return "bassoon";
case 71: return "clarinet";
case 72: return "flutes";
case 80: return "synth_lead";
case 88: return "synth_pad";
case DRUM_PROGRAM: return "drums";
default: return nullptr; // caller formats "program_"
}
}
// ---------------------------------------------------------------------------
// OpenNoteTracker — 1:1 port of events.py:96-231, the single state machine
// for both event decoding and prelude forcing.
// ---------------------------------------------------------------------------
struct NoteAction {
enum Kind { START, END, DRUM_HIT } kind;
int program; // rep program (unused for DRUM_HIT)
int pitch;
double time;
};
struct OpenNoteTracker {
// (program,pitch) -> onset, insertion-ordered like a Python dict
std::vector, double>> open;
double seek_time = 0.0, next_seek_time = -1.0; // <0 == None
int start_tick = 0, tick_state = 0;
int program = -1, velocity = -1; // -1 == None
bool in_prologue = true, skip_rest = false, chunk_started = false;
std::vector> tie_set;
bool open_has(std::pair key) const {
for (auto & e : open) if (e.first == key) return true;
return false;
}
void open_erase(std::pair key) {
for (size_t i = 0; i < open.size(); i++) {
if (open[i].first == key) { open.erase(open.begin() + (long) i); return; }
}
}
bool tie_has(std::pair key) const {
for (auto & e : tie_set) if (e == key) return true;
return false;
}
std::vector end_all(double time) {
std::vector a;
for (auto & e : open) a.push_back({ NoteAction::END, e.first.first, e.first.second, time });
open.clear();
return a;
}
std::vector feed_boundary(double seek, double next_seek /* <0 == None */) {
std::vector actions;
if (chunk_started && in_prologue) actions = end_all(seek_time);
seek_time = seek;
next_seek_time = next_seek;
start_tick = (int) llround(seek * FRAME_RATE);
tick_state = start_tick;
program = -1;
velocity = -1;
in_prologue = true;
skip_rest = false;
tie_set.clear();
chunk_started = true;
return actions;
}
std::vector feed(int token) {
Ev event = vocab_decode(token);
if (in_prologue) {
if (event.type == EV_TIE) {
in_prologue = false;
velocity = -1;
std::vector actions;
std::vector, double>> kept;
for (auto & e : open) {
if (tie_has(e.first)) kept.push_back(e);
else actions.push_back({ NoteAction::END, e.first.first, e.first.second, seek_time });
}
open = kept;
return actions;
}
if (event.type == EV_SHIFT) {
in_prologue = false;
skip_rest = true;
return end_all(seek_time);
}
if (event.type == EV_PROGRAM) {
program = event.value;
} else if (event.type == EV_PITCH && program >= 0) {
tie_set.push_back({ program, event.value });
}
return {};
}
if (skip_rest) return {};
if (event.type == EV_SHIFT) {
if (event.value > 0) tick_state = start_tick + event.value;
} else if (event.type == EV_PROGRAM) {
program = event.value;
} else if (event.type == EV_VELOCITY) {
velocity = event.value;
} else if (event.type == EV_DRUM) {
double time = (double) tick_state / FRAME_RATE;
if (next_seek_time < 0 || time < next_seek_time) {
return { { NoteAction::DRUM_HIT, DRUM_PROGRAM, event.value, time } };
}
} else if (event.type == EV_PITCH) {
if (program < 0 || velocity < 0) return {};
double time = (double) tick_state / FRAME_RATE;
if (next_seek_time >= 0 && time >= next_seek_time) return {};
std::pair key = { program, event.value };
std::vector actions;
if (open_has(key)) {
open_erase(key);
actions.push_back({ NoteAction::END, key.first, key.second, time });
}
if (velocity > 0) {
open.push_back({ key, time });
actions.push_back({ NoteAction::START, key.first, key.second, time });
}
return actions;
}
return {};
}
std::vector finish() {
if (chunk_started && in_prologue) return end_all(seek_time);
std::vector a;
for (auto & e : open) a.push_back({ NoteAction::END, e.first.first, e.first.second, e.second + MIN_NOTE_DUR });
open.clear();
return a;
}
// sorted (program,pitch) pairs currently held open (for tie prompts)
std::vector> open_keys() const {
std::vector> ks;
for (auto & e : open) ks.push_back(e.first);
std::sort(ks.begin(), ks.end());
return ks;
}
};
// mt3.py tie_section_token_ids: program token once per run of pitches, then tie
static std::vector tie_section_tokens(const std::vector> & open_keys) {
std::vector tokens;
int prog_state = -1;
for (auto & k : open_keys) {
if (k.first != prog_state) {
tokens.push_back(tok_program(k.first));
prog_state = k.first;
}
tokens.push_back(tok_pitch(k.second));
}
tokens.push_back(TOK_TIE);
return tokens;
}
// ---------------------------------------------------------------------------
// Note assembly + cleanup (tokenizer/notes.py validate/trim) + MIDI writer
// (utils/midi.py + note_event2midi — mido-compatible type-1 SMF)
// ---------------------------------------------------------------------------
struct MidiNote {
bool is_drum;
int program; // DRUM_PROGRAM for drums
double onset, offset;
int pitch;
};
static void sort_notes_vec(std::vector & notes) {
std::stable_sort(notes.begin(), notes.end(), [](const MidiNote & a, const MidiNote & b) {
if (a.onset != b.onset) return a.onset < b.onset;
if (a.is_drum != b.is_drum) return !a.is_drum;
if (a.program != b.program) return a.program < b.program;
if (a.pitch != b.pitch) return a.pitch < b.pitch;
return a.offset < b.offset;
});
}
static void validate_notes_fix(std::vector & notes) {
for (auto & n : notes) {
// matches validate_notes(fix=True): onset>offset -> max(offset, onset+0.01)
// which is always onset+0.01 in that branch; short non-drum notes padded
if (n.onset > n.offset) n.offset = n.onset + MIN_NOTE_DUR;
else if (!n.is_drum && n.offset - n.onset < 0.01) n.offset = n.onset + MIN_NOTE_DUR;
}
}
static std::vector trim_overlapping(std::vector notes) {
if (notes.size() <= 1) return notes;
// group by (program, pitch, is_drum); iterate groups in first-appearance order
std::vector out;
std::vector, bool>> seen;
for (auto & n : notes) {
std::pair, bool> ch = { { n.program, n.pitch }, n.is_drum };
bool dup = false;
for (auto & s : seen) if (s == ch) { dup = true; break; }
if (dup) continue;
seen.push_back(ch);
std::vector group;
for (auto & g : notes) {
if (g.program == n.program && g.pitch == n.pitch && g.is_drum == n.is_drum) group.push_back(g);
}
std::stable_sort(group.begin(), group.end(), [](const MidiNote & a, const MidiNote & b) { return a.onset < b.onset; });
for (size_t i = 1; i < group.size(); i++) {
if (group[i - 1].offset > group[i].onset) group[i - 1].offset = group[i].onset;
}
for (auto & g : group) if (g.onset < g.offset) out.push_back(g);
}
sort_notes_vec(out);
return out;
}
// SMF helpers (mido-compatible byte layout, no running status)
static void put_be32(std::vector & b, uint32_t v) {
b.push_back((uint8_t) (v >> 24)); b.push_back((uint8_t) (v >> 16));
b.push_back((uint8_t) (v >> 8)); b.push_back((uint8_t) v);
}
static void put_varlen(std::vector & b, uint32_t v) {
uint8_t buf[4];
int n = 0;
buf[n++] = (uint8_t) (v & 0x7f);
while (v >>= 7) buf[n++] = (uint8_t) ((v & 0x7f) | 0x80);
while (n--) b.push_back(buf[n]);
}
struct MidiEvent { // one channel/meta message at an absolute tick
double time;
bool is_drum;
int program; // track key
int velocity; // 1 = on, 0 = off (pre-writer semantics)
int pitch;
};
// events sorted like sort_note_events: (time, is_drum, program, velocity, pitch)
static void sort_midi_events(std::vector & evs) {
std::stable_sort(evs.begin(), evs.end(), [](const MidiEvent & a, const MidiEvent & b) {
if (a.time != b.time) return a.time < b.time;
if (a.is_drum != b.is_drum) return !a.is_drum;
if (a.program != b.program) return a.program < b.program;
if (a.velocity != b.velocity) return a.velocity < b.velocity;
return a.pitch < b.pitch;
});
}
// Serialize notes to a type-1 SMF (480 tpb, 120 bpm, velocity 100), matching
// note_event2midi: per-program named tracks, channels 0-8,10-15 by first
// appearance (overflow shares 15), drums on 9 with +0.01 s synthetic offs.
static std::vector write_midi(std::vector notes) {
const int TPB = 480;
const int TEMPO = 500000;
const int VEL = 100;
validate_notes_fix(notes);
notes = trim_overlapping(notes);
// note2note_event + drum offsets (writer adds them at time+0.01)
std::vector evs;
for (auto & n : notes) {
evs.push_back({ n.onset, n.is_drum, n.is_drum ? DRUM_PROGRAM : n.program, 1, n.pitch });
if (!n.is_drum) evs.push_back({ n.offset, false, n.program, 0, n.pitch });
}
sort_midi_events(evs);
{
std::vector drum_offs;
for (auto & e : evs) {
if (e.is_drum) drum_offs.push_back({ e.time + 0.01, true, DRUM_PROGRAM, 0, e.pitch });
}
for (auto & d : drum_offs) evs.push_back(d);
sort_midi_events(evs);
}
// per-program tracks
struct Track { std::vector bytes; int last_tick = 0; int channel = 0; };
std::vector track_order; // program keys, first appearance
std::vector