Skip to content

Native SMC: batched draws, autobatched potentials, llamppl removed - #149

Merged
shepardxia merged 193 commits into
mainfrom
shepard/speedups
Sep 26, 2026
Merged

shepardxia merged 193 commits into
mainfrom
shepard/speedups

Conversation

@shepardxia

@shepardxia shepardxia commented Jun 15, 2026 •

Copy link
Copy Markdown
Contributor

Removes llamppl, owns SMC natively, and batches everything a step does concurrently.

SMC

SMC validates and configures, smc_standard is the algorithm, SequenceModel is the particle, Sequences carries the results.

particles = await smc_standard(model, n_particles, ess_threshold=0.5)

smc_standard clones the template, gathers step() over live particles, records, tests ESS, and resamples by clone. One call is one SMC problem. Run several at once with asyncio.gather and their potential calls meet in the same batch windows.

SequenceModel.step() owns step semantics: draw, score, critic twist and settle, terminate. Samplers expose sample(context, draw=None) -> (token, logw, logp) and nothing else.

token_sampler.smc(...) keeps its signature. Breaking changes: SequenceModel and Importance leave the public sampler namespace, flatten_units moves to genlm.control.util, and llamppl is dropped as a dependency.

Batched draws

Concurrent draws across particles, across concurrent SMC runs, and inside unit-sampler loops meet in a per-event-loop window and resolve as one stacked reduction per (backend, device, vocab) group, off a single host readback per cohort. Rows stay on the backend's device until that readback.

draw_from(..., target=...) makes an entry an importance draw, sampled from the row and weighed under the target, served in the same readback.

set_draw_method selects the categorical sampler: gumbel_max (default), multinomial, inverse_cdf. All three draw the same distribution and differ in RNG consumption, which makes inverse_cdf the one to reach for in common-random-number experiments.

Autobatched potentials

SMC, DirectTokenSampler and AWRS wrap their potential seats in AutoBatchedPotential at construction, so concurrent per-particle asks reach the potential as one batched call. Pass autobatch=False to opt out. TrieSetSampler wraps only iter_potential, since the trie walk asks item_potential sequentially. autobatched() is one memoized door whose memo lives on the potential instance, so shared potentials share a window and the wrapper dies with its potential.

The old AsyncBatchLoop background task is deleted.

Reweighting

Tempered and Normalized are new:

p ** beta                      # scales every log weight; -inf survives any beta
p.normalize()                  # renormalizes each next-token row
(p ** (1 / tau)).normalize()   # p at temperature tau

Vocabulary-level sharing

Potential.build_tables with tables= skips the O(V) __init__. Coerced.build_trie with trie= shares one symbol trie across coercions over a vocabulary. live_logws lets a potential enumerate its live tokens, so batch_logw_next scatters a whole batch into one alloc_rows block; overriding alloc_rows places that block on a device. coerce(..., homomorphic=) declares whether f distributes over concatenation, which is the condition for taking the shared-prefix trie lane.

Testing

Behavioral tests only. New coverage for the new machinery: tests/test_draw_window.py, tests/sampler/test_smc_model.py, tests/potential/test_autobatch.py, tests/potential/test_reweight.py. The engine serving contract is asserted by the backend's GPU probe at row grain rather than through SMC statistics. pytest-xdist is added for parallel runs.

CI installs genlm-backend from requirements-ci.txt, which points at the co-developed branch. That file goes away when the backend releases and the floor in pyproject.toml moves up.

Lands with genlm-backend #74. Each side needs the other's branch.

shepardxia added 30 commits May 24, 2026 19:45
Hub owns the full per-token SMC algorithm (population, ESS, resample/fork,
log_ml); SlowDriver round-trips per token. All samplers collapse to one shared
transition. Vendors llamppl's resampler + SMCRecord; removes the llamppl dep.
Gated by tests/sampler/test_per_token_parity.py (18 cases, bit-parity vs the
prior smc_standard path). WindowDriver is a stub for the engine arms.

Co-authored by control-hub agent.
…om box ground truth

list(set(str tokens)) randomized vocab order per process (PYTHONHASHSEED unset),
making the Gumbel-max draw pick different tokens for the same noise. Original
smc_standard path had the same nondeterminism. Use list(dict.fromkeys(...)) for
hash-independent first-seen order; snapshot regenerated from the original path
on the box. Hub == original bit-for-bit under a pinned vocab.

Co-authored by control-hub agent.
Hub owns the draw (reuses fast_sample_lazyweights on engine logits) so window
matches slow by construction; ESS/resample/log_ml shared via extracted
_ess_crosses/_maybe_resample. Unconstrained DirectTokenSampler(llm) only;
constrained factor path stubbed + gated off (window_eligible). Gate-2a on box:
log_ml diff ~0, no systematic length/bias. Re-adds is_terminal_only/
supports_engine_native predicates.

Co-authored by window-integration agent.
…uction

Hub.draw reconstructs Product(llm, factor).logw_next in control V+1 space
(engine LM logits + factor.logw_next through the product index maps),
byte-identical to the slow path (max|ref-recon|=0). Factor's async logw_next
run via run_coroutine_threadsafe while run_window executes in a thread executor;
shape stays no-op. decompose_engine_target + window_eligible admit
DirectTokenSampler(llm*factor). Gate-2b on box: no systematic length/log_ml
bias (mean len gap +0.016 over 12 seeds). Speed 1.3-1.4x (factor logw_next
cost paid in both paths; window only saves LLM reprefill).

Co-authored by window-integration agent.
test_bytellm/test_llm spun up a fresh vLLM engine per test on the GPU box
(load_model_by_name defaults to vLLM there), cold-loading + recompiling and
OOMing at gpu_memory_utilization=0.9. Reuse a module-scoped engine at 0.3 util
with cleanup; skip test_mlx when mlx isn't importable (Linux box). No library
code changed. 13/13 test_bytellm on box.

Co-authored by window-integration agent.
…dedup imports

Behavior-preserving: shared ESS predicate between _ess_crosses (window pop-out)
and _maybe_resample (slow) so they decide identically; remove provably-dead
self.critic guards after the no-critic early return; module-level np/asyncio in
WindowDriver. Gate-1 18/18 on box after. shape/draw and vendored
resampling/smc_record untouched.

Co-authored by simplify-hub agent.
…t (drop self.backend)

Naming coherent with the library being 'control': Controller owns the SMC
algorithm and runs a StepLoop (per-token) or BurstLoop (engine). BurstLoop calls
self.llm.model.run_burst(...) directly, removing the self.backend plumbing.
window_eligible->can_burst, decompose_engine_target->split_engine_target.
Gate-1 25/25, gate-2 5/5 on box.
Split _compare into _run_slow/_run_window/_compare; cache the deterministic
StepLoop reference in gate2_snapshot.json (self-populated via GATE2_REGEN=1,
mirroring gate-1's parity_snapshot). Normal runs execute only the window and
compare against the snapshot (no-bias criteria, not byte-equality). 16 keys;
validated 5/5 on box (window vs fresh slow during regen).
… bit-identical)

Potentials become stateful (like the LLM's KV): Potential gains state0()/advance(
state, token)/logw_next_from_state(state). Trivial default carries the context
(byte-identical). WFSA overrides with its chart (what _consume already maintains),
so a consumer advances one symbol at a time instead of replaying context +
gathering 50k coroutines per step. BoolFSA booleanizes the stateful path too.
logw_next/prefix untouched (gate-1 ground truth intact). CPU test: logw_next_from_state
advanced through a context == logw_next(context) bit-for-bit (WFSA + BoolFSA,
incl. branching). DRY-ing the shared chart->weights step is a housekeeping task.
…-step gather)

Coerced threads the wrapped potential's state through the coerced byte-stream
(advance by f([token])'s symbols) and scores candidates from the carried chart
with a SYNC loop -- removing the asyncio.gather-over-the-vocab that profiling
showed is ~94% of logw_next's per-step cost (89.5/95ms at ~10k vocab; sync walk
~3.4x). WFSA gains sync prefix_logw/complete_logw (stateful analogs of prefix/
complete); BoolFSA booleanizes them. Falls back to recompute if the wrapped
potential isn't sync-stateful. CPU bit-identical to logw_next incl. dead-state
-inf (branching regex).
…te-opt)

Particle carries factor_state (cloned with context on resample, kept in lockstep);
BurstLoop seeds it (state0) at the first window and carries it across windows
(never re-derived); draw reads the factor via logw_next_from_state(state) -- sync
body, no per-step asyncio.gather over the vocab -- and advances it by each drawn
token. Slow path untouched (gate-1 18/18 local). Box: gate-2 5/5 parity vs the
cached snapshot, 5m39s vs ~14min (cache + stateful factor).
…window foundation)

hub.draw now loops + dispatches to a per-row strategy (w['draw_row']); the
DirectTokenSampler logic is extracted into Controller._draw_direct, with the
common LM-logits read and termination factored out. AWRS.sample is refactored to
delegate to a shared _run_rejection(logws, accept, target_logws) so the slow path
and (next) the window share the exact rejection algorithm + weight -- only the
logws source and accept form differ. Coerced gains sync prefix_logw/complete_logw
(delegating to the wrapped potential) for the stateful AWRS accept. Behavior-
preserving: gate-2 5/5 on box, slow AWRS 38 passed, gate-1 18/18.
…hub->controller

Samplers own their window draw: TokenSampler.supports_window()/window_draw();
DirectTokenSampler shapes the product, AWRS runs _run_rejection over the engine LM
logits with a stateful condition accept. Controller.draw calls unit_sampler.
window_draw polymorphically and banks -- no sampler-type dispatch. can_burst asks
supports_window(). Renamed module hub.py -> controller.py and swept every 'hub'
reference (vars, docstrings, imports, test helpers). Adds an AWRS gate-2 config.
NOT yet validated on the box (AWRS window path).
Replace the leaky raw `_window` dict passed to samplers with two explicit
objects: BurstContext (sampler-facing -- factor, product_logws(), factor_logws(),
run_sync(); no Controller runtime) and _Burst (Controller-private -- particles,
llm, eos_id, pop_out, ctx). Samplers' burst_draw now talks to a typed, named
interface, not string dict keys. Also rename the burst concept consistently
(WindowContext/_Window/window_draw/supports_window/_window/n_windows ->
Burst*/burst_draw/supports_burst/_burst/n_bursts) to match BurstLoop/run_burst/
can_burst. Slow path unchanged: gate-1 + stateful + seq sampler 30/30 local.
Burst path needs box re-validation.
Behavior-preserving cleanup from the multi-agent review of main..HEAD:
- controller.py: delete the unused control-side EngineControl Protocol (the
  backend owns the real contract; Controller.shape/draw are called directly),
  inline single-use _eos_idxs, drop a no-op bare return, fix module/comment
  drift (no "future" burst driver, no stale _advance_particle ref).
- token.py: hoist duplicated function-local imports (EndOfSequence,
  fast_sample_lazyweights) to module top; fix drifted base sample() signature.
- wfsa.py: BoolFSA booleanization was written 5x -> shared _booleanize/_bool.
- sequence.py: document the backend= arg; fix stale SequenceSampler example.
- coerce.py / wfsa.py: note the stateful-f homomorphism + shared-chart immutability.
- tests: conftest item_vocab list(set()) -> dict.fromkeys (PYTHONHASHSEED
  false-green); test_string_for_serialization drop unused fixture + assert;
  test_bytellm teardown warns instead of silently swallowing a failed cleanup.
Local gate-1 + stateful + seq sampler 30/30; burst path re-validated on box next.
The chart->weights tail was written twice (logw_next, logw_next_from_state) and
the prefix-weight normalizer three times (_prefix, prefix_logw,
logw_next_from_state). Extract _logw_next_from_chart (the shared tail) and
_chart_prefix_logw (the single dead-state/-nan boundary). Bit-identical by
construction: same chart, same log_ctx_w (raw, uncast). CPU bit-identical
test_stateful_potential + test_wfsa + gate-1 parity 36/36 local.
…ename

test_unconstrained_burst_vs_slow called _compare and discarded the result, so
the primary gate-2 case asserted nothing -- a real false-green in the acceptance
gate. Add the length + log_ml bounds (pure warm-KV residual: box shows log_ml
diff ~0, len diff 0.25). Also rename the leftover 'win' result vocabulary to
'burst' (vars + dict keys mean_len_win/win_log_ml -> *_burst) to match the
codebase's burst naming. Validated by the next box gate-2 run.
The context/number serialization was duplicated verbatim in
test_per_token_parity.py and _gen_parity_snapshot.py -- if they drift, the gate
silently compares against a mis-encoded reference. Move both into parity_cases
(the module both already import) so there is one definition. Drop the now-unused
EOS import from the gate. Gate-1 18/18 local.
- controller.draw: assert a particle that emits EOS to the engine has terminated
  in lockstep -- turns the slow(`is EOS`)/burst(`isinstance EndOfSequence`)
  divergence both review passes flagged into a loud failure instead of a silent
  length/log_ml gap. (Validated on box next regen.)
- test_engine_native: record result-affecting global config (model/prompt/eos)
  in the snapshot under __config__ and fail a load that mismatches it; a missing
  slow key now errors (run GATE2_REGEN=1) instead of silently self-comparing.
  Backward-compatible: legacy snapshots without __config__ still load. Wire the
  llm fixture + prompts to the recorded _EOS_BYTES/_PROMPT constants.
…ing-guard helpers

On-device [N,vocab]->[N,V+1] batch processing for the burst (one GPU pass, no
per-row numpy). Factor the EOS/non-EOS index-tensor caching into
_eos_index_tensors and the short-logits contiguity guard into
_verify_logit_padding, both shared by the per-row and batched paths so the
batched/burst path raises (not silently mis-folds) on non-contiguous added
tokens.
…twist

Replace the per-object Particle (+clone) with a columnar Population (logw/logp/
twist_amount/done/max_tokens_left as numpy arrays; contexts/states as lists) and
a thin Particle view onto row i -- single source of truth, no rebuilt
np.array([p.logw ...]) per step. ESS reads the arrays directly with a precomputed
log-threshold; resample is a columnar reindex; untwist is vectorized
(untwist_all/untwist_subset). draw() drops the dead sampling_metadata arg and
vectorizes the per-row untwist.
AWRS burst builds the rejection prep ON-DEVICE for the whole batch -- prune,
normalize (logps), Gumbel keys, logZ, and order=argsort -- then drives the
per-particle accept walk inline in the worker via _drive_sync (no event-loop
hop). The kernels take precomputed logps/order (slow path unchanged: order=None
draws+_TopK). Set burst gathers all particles' sample_set in ONE hop so the
backend trie batches across the population (was N serial run_sync hops).
MultiToken stays StepLoop-only (supports_burst=False).
Replace SMC.__call__'s backend= flag with a keyword-only accelerate="auto"|"off"|
"require" (True->auto, False->off). can_burst becomes burst_capability(controller)
-> BurstCapability(ok, reason); can_burst stays as a thin .ok wrapper. "require"
raises NotAcceleratable(reason); "auto" logs which path ran (reason on fallback).
Add SMC.acceleration_report() and thread accelerate through TokenSampler.smc.
SetTokenSampler.supports_burst()->False (Decision 2: don't silently deliver a
slowdown). genlm/control/__init__.py unchanged.
Replace the in-band EOS signaling with an out-of-band abort. _Burst.pop_out (which
forced a wasted t+1 forward to deliver a forced EOS) -> _Burst.abort_rows; draw flags
a row when its particle terminates (staggered) or all rows when ESS crosses, and
drain_aborts surfaces them for run_burst to abort_request. Every row in draw is now
live (aborted rows are dropped, never reappear), so the live-filtering and the
done/EOS-id forcing in _bank_burst_draw go away. Parity-identical (the t+1 forward
was always discarded); this is also the per-request abort primitive the interleave
and MultiToken pop-out need.
Replace the per-level logw_next(coerced_ctx + subtokens) replay with the carried
stateful interface: seed item_state through coerced_ctx once, advance one subtoken
per descent level, draw via logw_next_from_state. WFSA/BoolFSA item-potentials
advance their chart (kills the ~7.9s _logw_next_from_chart that dominated both Set
paths); generic potentials fall back to context-replay byte-identically. Linear
no-backtracking descent, so the state advances along the chosen path. Speeds the
slow AND burst Set paths; byte-exact (gate-1 18/18 + gate-2 9/9).
…ents)

Run B examples as ONE SMC population of B*N rows. Population gains a per-group
`group` column; ESS / resample / log_ml become per-group masked array ops and
per-row dispatch selects the group's sampler/critic -- the hot draw/bank paths
stay group-agnostic. Single-group (B=1) is byte-identical to before (gate-1 18/18).

Batched burst: one engine burst over B*N rows with IN-PLACE per-group resample --
when a group crosses ESS it is reindexed and its rows are aborted + re-added to the
live batch mid-burst (backend drain_adds), while the other groups keep decoding (no
engine drain). B=1 keeps the validated pop+relaunch path. Adds `batched_smc` /
`run_posterior_smc_batched` entry and a shared `_drive` driver-selection helper;
per-group prompt threading in the burst; on-device EOS gather in
_process_logw_next_batch.

Bundles the session's earlier engine-native work (stateful->context-scored factor,
AWRS probe-only, WFSA _consume LRU, Coerced homomorphism probe, shape removal,
burst critic batching).

Measured (Qwen2.5-Coder-7B, N=16, max_tokens=1024, B=4, StubCritic, ess=0, Blackwell,
vLLM 0.21): batched burst 13.0s == raw 1x64 decode ceiling; ~5x over parallel-estep,
~10x over vanilla/llamppl. gate-1 byte-identical; batched B>1 statistically no-bias
vs solo runs.
The await-0 window is collect_window in util.py; draw_from and
AutoBatchedPotential both call it. The autobatched memo moves onto the
potential instance (collectable, dies with its potential) and
to_autobatched resolves through it, so every route yields the same
wrapper and the same window. The view forwards is_terminal_only,
live_logws, and alloc_rows, and a contract test pins that every public
Potential method is defined on it. Dead code out: select(), the
iter_logws parameter, plot_bench's burst annotation, prof_entry's
hand-rolled run (bc.run_smc now, so --no-autobatch reaches SMC).
Draw flush reads back idx and one stacked float block; AWRS chains its
Gumbel keys in place.
Google-style sections, plain statements of contract, no history or
narration. Behavior untouched: every changed file's AST, docstrings
stripped, is identical to its previous version.
@shepardxia shepardxia changed the title Engine-native SMC speedups Native SMC: batched draws, autobatched potentials, llamppl removed Aug 17, 2026
Bare 'ndarray:' reads as a name to the google-style parser, so mkdocs strict
aborted on six warnings. Code is untouched; RNG consumption order still
matches llamppl.
The released backend has no per-request lora_name and pins transformers below
5, which mlx-lm now requires. Both jobs install the co-developed branch in the
same resolution as the project itself.
The acceleration section described a burst path that no longer exists.
engine_serving.md was a working note, not documentation. CLAUDE.md and local
dev artifacts are gitignored alongside .DS_Store.
tests/test_viz.py imports it. It used to arrive transitively under the old
transformers pin and does not under transformers 5.
engine_opts is vLLM-only and load_model_by_name falls back to HF without CUDA,
so passing it unconditionally errored every test in the module off-GPU.
The check is report-only, so the comment calls are best-effort; a 503 posting
the sticky comment was reddening it.
@shepardxia
shepardxia marked this pull request as ready for review August 17, 2026 21:00
@shepardxia
shepardxia requested review from ClementeP and samuki August 17, 2026 21:00
The harness is dev tooling for measuring this branch, not part of the library.
@shepardxia
shepardxia requested a review from vicky-xef August 24, 2026 16:15

@samuki samuki left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The refactor mostly looks good on a first glance. I put in a few questions and will run additional tests on the domains in genlm-eval.

Comment thread genlm/control/sampler/smc.py Outdated
Comment thread genlm/control/sampler/smc.py Outdated
Comment thread genlm/control/potential/autobatch.py
Comment thread genlm/control/util.py
Comment thread pyproject.toml Outdated
@@ -18,7 +18,6 @@ dependencies = [
"genlm-grammar>=0.2.0",
"genlm-backend>=0.2.0",

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We'll have to update this before the merge.

Comment thread genlm/control/util.py
Comment thread tests/potential/test_duplicate_byte_strings.py
collect_window handed its drained cohort back to a caller that could die
between queueing and flushing, stranding every co-caller's future. join_batch
now fails what a dying holder leaves behind, with a distinct BatchAbandoned
rather than the holder's own CancelledError -- that one, handed to a caller who
never asked for it, leaves their task cancelled and skips their
`except Exception`.

Both flushes gain the second arm: `except Exception` belongs to that group and
the flush continues, `except BaseException` resolves everything still owed and
re-raises. _flush_draws had only the second half and never re-raised;
autobatch._flush had only the first.

Every particle weight change funnels through _add_weight: NaN folds to -inf
where it happens rather than at the end of the run, and a +inf log-weight
raises, since that is what manufactures the NaN via W - w_sum. Without it
stratified, systematic and residual resampling collapse a whole population onto
one ancestor in silence. _unpack_particles' isnan coercion goes with it.

AWRS drew its rejection keys at float64 on the row's own device, which MPS does
not have at all; it now draws at the widest float the device carries.

terminate_when's no-correction contract moves to SMC.__call__ beside
max_tokens'. The three RNG streams are written down: np.random for resampling,
the global torch RNG for the pickers, a per-instance Generator for AWRS.
test_tokenize_roundtrip returns -- it pins the decode table as id-faithful over
a duplicate-byte-string vocab. test_bytellm's teardown stops running
asyncio.run over a sync cleanup, which threw into a warning every time and
would have masked a real teardown failure.

543 passed local, 540 passed serial on l40s.
Comment thread genlm/control/util.py Outdated

@vicky-xef vicky-xef left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We should also check docs/ folder and README.md files before merging

Comment thread genlm/control/sampler/token.py
Comment thread genlm/control/util.py Outdated
Comment thread genlm/control/util.py Outdated
Comment thread genlm/control/util.py Outdated
The weights docstrings claimed numpy, but a row off a vLLM engine arrives as a
torch tensor and the constructor has always taken both.

The batching counters go with them: `take_batch_stats` rebound the module
global, so `from ... import batch_stats` in autobatch.py kept incrementing an
orphan and its counts never reached a caller.

The tutorial still taught a manual `to_autobatched()` on the critic and timed
it against a run that was already batched, since SMC wraps the critic unless
`autobatch=False`.
`PotentialTests` gains `assert_contract`, so a potential's three properties are
asserted by one call instead of a hand-wired triple at 58 sites.

Three of four resamplers had never executed: nothing passed `resampling_method`
and nothing called them directly. `test_resampling.py` sweeps all four, and
pins the floor/ceil bound that is the reason to prefer systematic or stratified
over multinomial. The distribution check likewise ran only `gumbel_max`; it now
sweeps `DRAW_METHODS`, which retires a round-trip test that asserted only that
indices were in range.

Removes tests subsumed by a stronger sibling, tests whose only assertion was
incidental to another test's subject, and duplicate parametrize entries.
A fresh `torch.Generator` starts from a fixed state, so two unseeded samplers
drew the same rejection keys. `seed=None` now means the global torch RNG:
independent per instance, and reproducible under `torch.manual_seed` like the
rest of the draw path. A seed still pins its own stream.

The SWOR tracer took its mass from the trie in float32, where a branch worth
1e-50 is annihilated by subtraction from a sibling worth 1e-5 -- the walk hit
zero mass and stopped one token short of the target, which read as a sampler
bug. Masses are float64 now.
Add Potential.tables and use it in every wrapper; PromptedLLM goes through
Potential.__init__ instead of copying its fields. WFSA.spawn keeps
cache_maxsize and BoolFSA accepts it. Rename picker -> draw method,
live_logws -> sparse_logw_next, and the batch vocabulary to plain words.
Trim every docstring and comment the branch added to main's register,
restoring main's text where behaviour did not change. Drop dead code and a
duplicate test; restore the single-token unit-sampler case. Docs: the
performance page covers backend choice, default autobatching and concurrent
SMC runs; potentials and samplers cover Tempered, Normalized, homomorphic=,
lora_name= and set_draw_method.
@shepardxia
shepardxia merged commit 490d1d9 into main Sep 26, 2026
5 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants