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: Protocol

Anything that can suggest the next slot to query.

Implementations must be synchronous from the policy’s point of view.

prefer(observed, available, profile)#

Return a preferred slot from available (or None).

Return type:

str | None

Parameters:
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: object

An ActionSupervisor backed by the TRL unified LLM interface.

Parameters:
async aprefer(observed, available, profile)#

Async: ask the LLM for the most informative remaining slot.

Return type:

str | None

Parameters:
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.

Return type:

str | None

Parameters:
property token_usage: TokenCounter#

Accumulated token usage across supervisor calls.

pydantic model NextActionPreference#

Bases: BaseModel

Structured 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: object

The outcome of running the policy on a single sample.

Parameters:
probability#

Classifier success probability from the final state.

prediction#

Binary label (1 if probability >= predict_threshold).

slots_used#

Distinct information slots queried (in first-query order).

decision_path#

Ordered actions taken, each a slot name or "stop".

decision_path: List[str]#
prediction: int#
probability: float#
slots_used: List[str]#
class VerifiableRL(slots, config=None, supervisor=None, device='cpu', random_state=None)#

Bases: object

Adaptive, sequential information-gathering binary classifier.

Parameters:
  • slots (Union[Sequence[str], Mapping[str, int]]) – Either an ordered sequence of slot names (feature dimensions are inferred at fit) or a mapping of name -> dimension (lets state_dim be known before fitting and enables from_state_dicts()).

  • config (VerifiableRLConfig | None) – Algorithm configuration. Defaults to VerifiableRLConfig().

  • supervisor (ActionSupervisor | None) – Optional ActionSupervisor consulted when the policy is undecided. Requires per-sample profiles at 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.pt format. slots must be a name -> dimension mapping (the dimensions are required to rebuild the networks) whose total matches the checkpoint’s state_dim and whose count matches action_dim - 1.

Return type:

VerifiableRL

Parameters:
classmethod load(path, *, supervisor=None, device=None)#

Load a model saved with save().

Return type:

VerifiableRL

Parameters:
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 (dict name -> (n, dim) array, or a list of arrays aligned to slots).

  • y (Any) – Binary labels (length n, 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,)).

Return type:

ndarray

Parameters:
predict_paths(X, *, profiles=None)#

Run the policy on each sample and return full decision traces.

Return type:

List[QueryResult]

Parameters:
predict_proba(X, *, profiles=None)#

Return per-sample success probabilities (shape (n,)).

Return type:

ndarray

Parameters:
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.py curriculum: early epochs reveal all slots, later epochs reveal progressively fewer, so the classifier learns to predict from partial information. The returned state_dict is meant to warm-start training via fit(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 (dict name -> (n, dim) or list aligned to slots).

  • y (Any) – Binary labels (length n, 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 (frac should 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’s random_state).

Return type:

Dict[str, Tensor]

Returns:

A classifier state_dict of CPU tensors.

save(path)#

Save weights + config to a directory.

Return type:

None

Parameters:

path (str | PathLike[str])

property action_dim: int#

one per slot plus STOP.

Type:

Number of actions

property is_fitted: bool#

Whether the model has trained or loaded weights.

property state_dim: int | None#

Length of the state vector (None until slot dims are known).

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: object

Configuration 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 in fit.

  • 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. 0 disables.

  • 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. 0 makes 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)

clf_batch: int = 512#
clf_hidden: int = 256#
clf_lr: float = 1e-05#
clf_replay_max: int = 40000#
clf_sample: int = 4000#
clf_target_update_every: int = 5#
clf_threshold: float = 0.3#
eps: float = 0.1#
freeze_clf_updates: int = 10#
grad_clip: float = 5.0#
greedy: bool = True#
llm_bias: float = 0.05#
max_depth: int = 4#
max_steps: int = 5#
min_queries: int = 3#
n_iterations: int = 1#
n_rollouts: int = 10#
policy_batch: int = 128#
policy_hidden: int = 512#
policy_lr: float = 5e-05#
policy_replay_max: int = 20000#
policy_sample: int = 2000#
predict_min_queries: int = 1#
predict_threshold: float = 0.3#
repeat_penalty: float = -5.0#
reward_fn: float = -0.25#
reward_fp: float = -16.0#
reward_tn: float = 0.0#
reward_tp: float = 4.0#
step_penalty: float = -0.1#
tau_info: float = 1.0#
tau_stop: float = 4.0#
train_epochs: int = 1#
uncertain_delta: float = 0.01#
update_every: int = 25#