Verifiable RL#
Module contents#
Verifiable RL.
An adaptive information-gathering binary classifier: a learned policy network decides, slot by slot, which information to reveal next (or to STOP), guided by Monte-Carlo tree-search targets and optional weak LLM supervision; a classifier then predicts the label from the accumulated partial state.
- class ActionSupervisor(*args, **kwargs)#
Bases:
ProtocolAnything that can suggest the next slot to query.
Implementations must be synchronous from the policy’s point of view.
- class LLMActionSupervisor(llmc, *, instructions='You are a decision-support module for an information-gathering agent. Choose the ONE remaining information slot most likely to change the final success/failure decision. Do not suggest already-observed slots. Prefer the slot with the highest marginal information value. Return at most one slot as a weak reference; if nothing stands out, return an empty list. Return structured output only.', temperature=0.0, llm_semaphore_limit=3, _llm=None)#
Bases:
objectAn
ActionSupervisorbacked by the TRL unified LLM interface.- Parameters:
llmc (
List[Union[AnthropicChoice,GoogleChoice,OpenAIChoice,XAIChoice,AnthropicChoiceDict,GoogleChoiceDict,OpenAIChoiceDict,XAIChoiceDict]]) – LLM choices in priority order.instructions (
str) – System instructions for the preference call.temperature (
float) – Sampling temperature (default deterministic).llm_semaphore_limit (
int) – Max concurrent LLM calls._llm (
Any) – LLM instance for dependency injection (testing). If None, uses the globalllmsingleton.
- async aprefer(observed, available, profile)#
Async: ask the LLM for the most informative remaining slot.
- prefer(observed, available, profile)#
Sync wrapper around
aprefer(), with per-state caching.The coroutine runs in an isolated worker thread so the caller’s event loop (and the process-wide default loop) is never touched; this keeps the synchronous training/prediction path safe to call from anywhere.
- property token_usage: TokenCounter#
Accumulated token usage across supervisor calls.
- pydantic model NextActionPreference#
Bases:
BaseModelStructured weak preference returned by the LLM action supervisor.
- prefer#
Zero or one preferred slot name. The supervisor only nudges the policy; an empty list means “no clear preference”.
- field prefer: List[str] [Optional]#
- class QueryResult(probability, prediction, slots_used, decision_path)#
Bases:
objectThe outcome of running the policy on a single sample.
- probability#
Classifier success probability from the final state.
- prediction#
Binary label (
1ifprobability >= predict_threshold).
- slots_used#
Distinct information slots queried (in first-query order).
- decision_path#
Ordered actions taken, each a slot name or
"stop".
- class VerifiableRL(slots, config=None, supervisor=None, device='cpu', random_state=None)#
Bases:
objectAdaptive, sequential information-gathering binary classifier.
- Parameters:
slots (
Union[Sequence[str],Mapping[str,int]]) – Either an ordered sequence of slot names (feature dimensions are inferred atfit) or a mapping ofname -> dimension(letsstate_dimbe known before fitting and enablesfrom_state_dicts()).config (
VerifiableRLConfig|None) – Algorithm configuration. Defaults toVerifiableRLConfig().supervisor (
ActionSupervisor|None) – OptionalActionSupervisorconsulted when the policy is undecided. Requires per-sampleprofilesat fit/predict time.device (
str) – Torch device string (e.g."cpu"or"cuda").random_state (
int|None) – Seed for reproducible network init and rollouts.
- classmethod from_state_dicts(path, slots, *, config=None, supervisor=None, device='cpu')#
Load a raw checkpoint of
policy_state_dict+clf_state_dict.This consumes the standalone
runs/model_*/final_model.ptformat.slotsmust be aname -> dimensionmapping (the dimensions are required to rebuild the networks) whose total matches the checkpoint’sstate_dimand whose count matchesaction_dim - 1.- Return type:
- Parameters:
config (VerifiableRLConfig | None)
supervisor (ActionSupervisor | None)
device (str)
- classmethod load(path, *, supervisor=None, device=None)#
Load a model saved with
save().- Return type:
- Parameters:
supervisor (ActionSupervisor | None)
device (str | None)
- fit(X, y, *, profiles=None, pretrained_classifier=None)#
Train the policy and classifier on partially observable episodes.
Mirrors the standalone two-network training loop: an optional classifier warm-start, a freeze period during which only the policy trains against the (fixed) classifier, and a slowly-synced target classifier that the MCTS rollouts query so the policy’s value targets stay stable.
- Parameters:
X (
Union[Mapping[str,ndarray],Sequence[ndarray]]) – Slot features (dictname -> (n, dim)array, or a list of arrays aligned toslots).y (
Any) – Binary labels (lengthn, values in{0, 1}).profiles (
Optional[Sequence[str]]) – Optional per-sample text used by the LLM supervisor.pretrained_classifier (
Union[str,PathLike[str],Mapping[str,Any],None]) – Optional warm-start for the classifier — a path to a checkpoint ({"state_dict": ...},{"clf_state_dict": ...}, or a raw state dict) or an in-memory state dict. Strongly recommended for faithful results: the MCTS reward depends on the classifier, so starting it from a pretrained model (rather than random) is what lets the policy learn useful queries (matches the reference implementation).
- Return type:
Self- Returns:
self.
- predict(X, *, profiles=None)#
Return per-sample binary predictions (shape
(n,)).
- predict_paths(X, *, profiles=None)#
Run the policy on each sample and return full decision traces.
- predict_proba(X, *, profiles=None)#
Return per-sample success probabilities (shape
(n,)).
- pretrain_classifier(X, y, *, epochs=50, lr=1e-05, batch_size=16, curriculum=None, random_state=None)#
Pretrain a classifier on randomly-masked partial states.
Faithfully mirrors the standalone
pretrain_classifier.pycurriculum: early epochs reveal all slots, later epochs reveal progressively fewer, so the classifier learns to predict from partial information. The returnedstate_dictis meant to warm-start training viafit(pretrained_classifier=...)— which is what makes the MCTS reward (and therefore the learned policy) meaningful.This does not mutate the model; it returns weights to pass to
fit. For a leakage-free cross-validation, call this on each fold’s training data only.- Parameters:
X (
Union[Mapping[str,ndarray],Sequence[ndarray]]) – Slot features (dictname -> (n, dim)or list aligned to slots).y (
Any) – Binary labels (lengthn, values in{0, 1}).epochs (
int) – Number of passes over the data.lr (
float) – Adam learning rate.batch_size (
int) – Mini-batch size.curriculum (
Optional[Sequence[Mapping[str,Any]]]) – Optional list of{"min_k", "max_k", "frac"}phases (fracshould sum to ~1). Defaults to a full->medium->light schedule derived from the number of slots.random_state (
int|None) – Seed for masking/shuffling (falls back to the instance’srandom_state).
- Return type:
- Returns:
A classifier
state_dictof CPU tensors.
- save(path)#
Save weights + config to a directory.
- class VerifiableRLConfig(policy_hidden=512, clf_hidden=256, max_steps=5, max_depth=4, n_rollouts=10, min_queries=3, predict_min_queries=1, reward_tp=4.0, reward_fp=-16.0, reward_tn=0.0, reward_fn=-0.25, step_penalty=-0.1, repeat_penalty=-5.0, clf_threshold=0.3, tau_info=1.0, tau_stop=4.0, n_iterations=1, update_every=25, policy_lr=5e-05, clf_lr=1e-05, policy_batch=128, clf_batch=512, train_epochs=1, grad_clip=5.0, eps=0.1, policy_replay_max=20000, clf_replay_max=40000, policy_sample=2000, clf_sample=4000, freeze_clf_updates=10, clf_target_update_every=5, uncertain_delta=0.01, llm_bias=0.05, predict_threshold=0.3, greedy=True)#
Bases:
objectConfiguration for
VerifiableRL.Defaults mirror the standalone VCBench reference implementation (
runs/model_*checkpoints were trained with these network sizes and reward values).- Parameters:
policy_hidden (
int) – Hidden width of the PolicyNet MLP.clf_hidden (
int) – Hidden width of the Classifier MLP.max_steps (
int) – Max actions per episode during training/prediction.max_depth (
int) – Max rollout depth inside MCTS.n_rollouts (
int) – Monte-Carlo rollouts per root action.min_queries (
int) – Slots that must be revealed before STOP is allowed when building policy targets.predict_min_queries (
int) – Same guard, applied at prediction time.reward_tp/reward_fp/reward_tn/reward_fn – Asymmetric terminal rewards.
step_penalty (
float) – Per-step penalty added to a terminal reward.repeat_penalty (
float) – Penalty for re-querying an already-observed slot.clf_threshold (
float) – Probability threshold used inside the reward function.tau_info (
float) – Temperature for the info-action softmax target.tau_stop (
float) – Temperature for the STOP-action sigmoid target.n_iterations (
int) – Passes over the training set infit.update_every (
int) – Train the networks every this-many episodes.policy_lr/clf_lr – Adam learning rates.
policy_batch/clf_batch – Mini-batch sizes.
train_epochs (
int) – Optimisation epochs per network update.grad_clip (
float) – Gradient-norm clip (None/0 disables).eps (
float) – Epsilon for epsilon-greedy exploration during training.policy_replay_max/clf_replay_max – Replay-buffer caps.
policy_sample/clf_sample – Samples drawn per update (
0= use all).freeze_clf_updates (
int) – Number of initial network updates during which the classifier is held fixed (only the policy trains) so the policy can first adapt to a stable (warm-started) classifier.0disables.clf_target_update_every (
int) – Sync the frozen target classifier (the one the MCTS rollouts/rewards query) to the live classifier every this many updates.0makes rollouts use the live classifier directly.uncertain_delta (
float) – Top-2 info-prob gap below which the supervisor is asked.llm_bias (
float) – Probability mass added to the supervisor’s preferred slot.predict_threshold (
float) – Probability threshold for the final binary label.greedy (
bool) – If True, take argmax actions at prediction time.reward_tp (float)
reward_fp (float)
reward_tn (float)
reward_fn (float)
policy_lr (float)
clf_lr (float)
policy_batch (int)
clf_batch (int)
policy_replay_max (int)
clf_replay_max (int)
policy_sample (int)
clf_sample (int)