Native SMC: batched draws, autobatched potentials, llamppl removed - #149
Merged
Merged
Conversation
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).
…straint alpha/json)
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.
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
marked this pull request as ready for review
August 17, 2026 21:00
The harness is dev tooling for measuring this branch, not part of the library.
samuki
requested changes
Aug 27, 2026
samuki
left a comment
Member
There was a problem hiding this comment.
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.
| @@ -18,7 +18,6 @@ dependencies = [ | |||
| "genlm-grammar>=0.2.0", | |||
| "genlm-backend>=0.2.0", | |||
Member
There was a problem hiding this comment.
We'll have to update this before the merge.
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.
samuki
reviewed
Sep 17, 2026
vicky-xef
reviewed
Sep 21, 2026
vicky-xef
left a comment
Contributor
There was a problem hiding this comment.
We should also check docs/ folder and README.md files before merging
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.
samuki
approved these changes
Sep 25, 2026
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.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Removes llamppl, owns SMC natively, and batches everything a step does concurrently.
SMC
SMCvalidates and configures,smc_standardis the algorithm,SequenceModelis the particle,Sequencescarries the results.smc_standardclones the template, gathersstep()over live particles, records, tests ESS, and resamples by clone. One call is one SMC problem. Run several at once withasyncio.gatherand their potential calls meet in the same batch windows.SequenceModel.step()owns step semantics: draw, score, critic twist and settle, terminate. Samplers exposesample(context, draw=None) -> (token, logw, logp)and nothing else.token_sampler.smc(...)keeps its signature. Breaking changes:SequenceModelandImportanceleave the publicsamplernamespace,flatten_unitsmoves togenlm.control.util, andllampplis 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_methodselects the categorical sampler:gumbel_max(default),multinomial,inverse_cdf. All three draw the same distribution and differ in RNG consumption, which makesinverse_cdfthe one to reach for in common-random-number experiments.Autobatched potentials
SMC,DirectTokenSamplerandAWRSwrap their potential seats inAutoBatchedPotentialat construction, so concurrent per-particle asks reach the potential as one batched call. Passautobatch=Falseto opt out.TrieSetSamplerwraps onlyiter_potential, since the trie walk asksitem_potentialsequentially.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
AsyncBatchLoopbackground task is deleted.Reweighting
TemperedandNormalizedare new:Vocabulary-level sharing
Potential.build_tableswithtables=skips the O(V)__init__.Coerced.build_triewithtrie=shares one symbol trie across coercions over a vocabulary.live_logwslets a potential enumerate its live tokens, sobatch_logw_nextscatters a whole batch into onealloc_rowsblock; overridingalloc_rowsplaces that block on a device.coerce(..., homomorphic=)declares whetherfdistributes 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-xdistis 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 inpyproject.tomlmoves up.Lands with genlm-backend #74. Each side needs the other's branch.