// Sortformer diarization engine: 16 kHz mono WAV in, speaker turns out. // // Kept deliberately dumb, exactly like scripts/diarize-sherpa.py: read audio, // run the model, print raw turns as JSON on stdout. All policy (filenames, // provenance shape, atomic writes) lives in scripts/diarize.mjs, and audio // decoding lives in scripts/diarize-sortformer.mjs, so this file is only ever // "run the model". // // stdout is JSON ONLY. Anything human goes to stderr. // // WHY THIS FILE EXISTS RATHER THAN AN UPSTREAM BINARY: openresearchtools/engine // ships `llama-realtime-smoke`, a PARITY tool. It constructs the backend with // capture_debug=true (retaining every intermediate matrix for every step), it // requires a directory of PyTorch reference fixtures we do not have, and its // JSON dump drops the event flags — which is fatal, because 2033 of the 2099 // events it emits for a 12.8-minute file are PREVIEW re-emissions of spans that // are still growing. Only the 66 non-preview events are the answer. This driver // is that tool with the debug capture off, the fixtures gone, and the preview // events filtered. // // WHY NO MERGE PASS: the non-preview events are already the postprocessed, // disjoint spans — running an interval union over them is a no-op (66 in, 66 // out, verified). Unlike the sherpa engine there is no clustering step and no // threshold: Sortformer is end-to-end, which is the whole reason it does not // over-split. // // WHY RAW PCM ON STDIN IS THE PRIMARY INPUT: this corpus has 8-hour VODs. A // decoded 16 kHz mono WAV of one is ~900 MB on a disk that is 98% full, and // buffering it as float32 is ~460 MB resident. Streaming ffmpeg's output straight // into push_audio() costs neither: the model's own state is a fixed-size speaker // cache, so memory becomes O(1) in duration rather than O(n). That is the single // biggest advantage over the sherpa engine, whose O(n^2) clustering matrix is // what forced the whole windowing/centroid-reclustering design in // diarize-sherpa.py. --audio-wav is kept for direct CLI use and testing. #include "backend-factory.h" #include "sortformer/sortformer-backend.h" #include "stream-manager.h" #include #include #include #include #include #include #include #include #include #include #include #include #include // POSIX only, which is all scripts/build-sortformer.sh targets. The upstream // engine builds on Windows too; this driver deliberately does not try to. #include #include #include namespace { // 16-bit PCM mono, which is what `ffmpeg -ac 1 -ar 16000 -c:a pcm_s16le` makes. // Deliberately not a general WAV reader: the adapter always hands us that exact // format, and quietly accepting anything else would mean silently diarizing // resampled-wrong audio. std::vector load_wav_s16_mono(const std::string & path, uint32_t & sample_rate) { std::ifstream f(path, std::ios::binary); if (!f) throw std::runtime_error("cannot open wav: " + path); char riff[12]; f.read(riff, 12); if (std::strncmp(riff, "RIFF", 4) != 0 || std::strncmp(riff + 8, "WAVE", 4) != 0) throw std::runtime_error("not a RIFF/WAVE file: " + path); uint16_t channels = 0, bits = 0; while (f) { char id[4]; uint32_t sz = 0; f.read(id, 4); f.read(reinterpret_cast(&sz), 4); if (!f) break; if (std::strncmp(id, "fmt ", 4) == 0) { std::vector fmt(sz); f.read(fmt.data(), sz); if (sz >= 16) { std::memcpy(&channels, fmt.data() + 2, 2); std::memcpy(&sample_rate, fmt.data() + 4, 4); std::memcpy(&bits, fmt.data() + 14, 2); } } else if (std::strncmp(id, "data", 4) == 0) { if (channels != 1 || bits != 16) throw std::runtime_error("expected 16-bit mono wav, got " + std::to_string(channels) + "ch/" + std::to_string(bits) + "bit"); const size_t n = sz / 2; std::vector pcm(n); f.read(reinterpret_cast(pcm.data()), sz); std::vector out(n); for (size_t i = 0; i < n; ++i) out[i] = static_cast(pcm[i]) / 32768.0f; return out; } else { f.seekg(sz + (sz & 1), std::ios::cur); // chunks are word-aligned } } throw std::runtime_error("no data chunk in wav: " + path); } // The sortformer path never sets a thread count (only voxtral's runtime does), // so the CPU backend runs at ggml's default of 4. On this box 6 threads is the // optimum (1305 s/audio-hour) and 8 is WORSE than 4 through oversubscription, so // this has to be reachable rather than left to the default. void set_backend_n_threads(ggml_backend_t backend, int n_threads) { if (backend == nullptr) return; ggml_backend_dev_t dev = ggml_backend_get_device(backend); ggml_backend_reg_t reg = dev ? ggml_backend_dev_backend_reg(dev) : nullptr; if (reg == nullptr) return; auto * fn = (ggml_backend_set_n_threads_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_set_n_threads"); if (fn != nullptr) fn(backend, n_threads); } // Stream s16le mono PCM from stdin straight into the session. Never holds more // than one block, so an 8-hour VOD costs the same resident memory as a 6-minute // clip. Returns the number of samples fed. size_t feed_raw_stdin(llama::realtime::stream_manager & mgr, int64_t sid, uint32_t sample_rate, size_t block_samples) { std::vector pcm(block_samples); std::vector block(block_samples); size_t total = 0; // FORCE THE PIPE BACK TO BLOCKING. Node sets the pipes it hands a child to // non-blocking, so a plain read() returns -1/EAGAIN long before EOF and the // stream looks like an I/O error. It works from a shell and fails under the // app, which is exactly the sort of difference that gets found in production // rather than in a test — so it is fixed here, in the binary, rather than // being made the caller's problem. const int flags = fcntl(STDIN_FILENO, F_GETFL, 0); if (flags != -1 && (flags & O_NONBLOCK)) fcntl(STDIN_FILENO, F_SETFL, flags & ~O_NONBLOCK); // read(2) rather than fread: it lets EINTR be retried without the ambiguity // of a short fread, and a partial read is normal on a pipe rather than an // error to distinguish from EOF. while (true) { size_t filled = 0; const size_t want = block_samples * sizeof(int16_t); auto * buf = reinterpret_cast(pcm.data()); while (filled < want) { const ssize_t n = ::read(STDIN_FILENO, buf + filled, want - filled); if (n == 0) break; // EOF if (n < 0) { if (errno == EINTR) continue; throw std::runtime_error(std::string("read error on stdin: ") + std::strerror(errno)); } filled += static_cast(n); } const size_t got = filled / sizeof(int16_t); if (got == 0) break; for (size_t i = 0; i < got; ++i) block[i] = static_cast(pcm[i]) / 32768.0f; mgr.push_audio(sid, block.data(), got, sample_rate); total += got; if (filled < want) break; // short read means EOF } return total; } void usage(const char * argv0) { std::cerr << "usage: " << argv0 << " --model (--audio-raw - | --audio-wav )\n" << " [--sample-rate 16000] [--backend Vulkan0|CPU]\n" << " [--threads N] [--feed-ms N]\n\n" << "--audio-raw - read s16le mono PCM from stdin (streaming, O(1) memory)\n" << "--audio-wav read a 16 kHz mono 16-bit WAV file\n\n" << "Prints {\"turns\":[{start,end,speaker}],\"audioSeconds\":N,\"version\":\"...\"}\n" << "on stdout. Progress and errors go to stderr.\n"; } } // namespace int main(int argc, char ** argv) { try { std::string gguf, wav, raw, backend = "Vulkan0"; uint32_t sr = 16000; int n_threads = 0; double feed_ms = 100.0; for (int i = 1; i < argc; ++i) { const std::string a = argv[i]; if (a == "--help" || a == "-h") { usage(argv[0]); return 0; } if (i + 1 >= argc) throw std::invalid_argument("missing value for " + a); if (a == "--model") gguf = argv[++i]; else if (a == "--audio-wav") wav = argv[++i]; else if (a == "--audio-raw") raw = argv[++i]; else if (a == "--sample-rate") sr = static_cast(std::stoul(argv[++i])); else if (a == "--backend") backend = argv[++i]; else if (a == "--threads") n_threads = std::stoi(argv[++i]); else if (a == "--feed-ms") feed_ms = std::stod(argv[++i]); else throw std::invalid_argument("unknown argument: " + a); } if (gguf.empty() || (wav.empty() == raw.empty())) { usage(argv[0]); return 2; } if (!raw.empty() && raw != "-") throw std::invalid_argument("--audio-raw only supports '-' (stdin)"); // Load the model BEFORE reading stdin, so a bad model path fails before // ffmpeg has decoded anything. capture_debug=false — see the header. auto be = std::make_unique(gguf, backend, false); auto * be_ptr = be.get(); if (n_threads > 0) set_backend_n_threads(be_ptr->model().backend(), n_threads); llama::realtime::stream_manager mgr; const int64_t sid = mgr.create_session(std::move(be)); std::cerr << "sortformer: backend=" << be_ptr->backend_name() << " input=" << (raw.empty() ? wav : "stdin") << "\n"; const auto t0 = std::chrono::steady_clock::now(); size_t n_samples = 0; if (!raw.empty()) { if (sr == 0) throw std::runtime_error("--sample-rate must be non-zero"); n_samples = feed_raw_stdin( mgr, sid, sr, std::max(1, static_cast((feed_ms / 1000.0) * sr))); } else { const auto audio = load_wav_s16_mono(wav, sr); if (sr == 0) throw std::runtime_error("wav reports a zero sample rate"); const size_t feed = std::max(1, static_cast((feed_ms / 1000.0) * sr)); for (size_t off = 0; off < audio.size(); off += feed) { const size_t n = std::min(feed, audio.size() - off); mgr.push_audio(sid, audio.data() + static_cast(off), n, sr); } n_samples = audio.size(); } mgr.flush_session(sid); const double infer_sec = std::chrono::duration(std::chrono::steady_clock::now() - t0).count(); const double audio_sec = static_cast(n_samples) / static_cast(sr); // PREVIEW EVENTS ARE NOT THE ANSWER. A streaming span is re-emitted every // time it grows; only the final, postprocessed commits carry the result. const auto events = mgr.drain_events(sid, 0); std::ostringstream turns; size_t n_turns = 0; for (const auto & e : events) { if (e.type != llama::realtime::event_type::speaker_span_commit) continue; if (e.flags & llama::realtime::event_flag_preview) continue; if (e.end_sec <= e.begin_sec) continue; if (n_turns++) turns << ","; turns << "{\"start\":" << e.begin_sec << ",\"end\":" << e.end_sec << ",\"speaker\":" << e.speaker_id << "}"; } std::cerr << "sortformer: " << n_turns << " turns in " << std::fixed << std::setprecision(1) << infer_sec << "s (" << std::setprecision(2) << (audio_sec / infer_sec) << "x realtime)\n"; std::cout << std::setprecision(6) << "{\"turns\":[" << turns.str() << "]" << ",\"audioSeconds\":" << audio_sec << ",\"version\":\"" << be_ptr->backend_name() << "\"}\n"; return 0; } catch (const std::exception & e) { std::cerr << "sortformer: error: " << e.what() << "\n"; return 1; } }