Policy Induction#
Module contents#
Policy Induction.
An interpretable ensemble framework for binary classification based on policies generated by LLMs.
- pydantic model Policies#
Bases:
BaseModel- field policies: List[str] [Required]#
The list of generated policies.
- class PolicyInduction(gen_llmc, predict_llmc=None, config=None, gen_temperature=1.0, predict_temperature=0.0, llm_semaphore_limit=3, max_policy_length=20, class_ratio=(1.0, 1.0), max_samples_as_context=10, p_predict_update_interval=10, save_path=None, name=None, random_state=0, confirm_requests=True)#
Bases:
objectInterpretable ensemble binary classifier.
Induces natural-language policies from labeled data via LLM, scores each policy against every sample, then trains a logistic regression to find the optimal weighted combination for YES/NO prediction.
- Parameters:
gen_llmc (
List[Union[AnthropicChoice,GoogleChoice,OpenAIChoice,XAIChoice,AnthropicChoiceDict,GoogleChoiceDict,OpenAIChoiceDict,XAIChoiceDict]]) – LLMs for policy generation, in priority order.predict_llmc (
Optional[List[Union[AnthropicChoice,GoogleChoice,OpenAIChoice,XAIChoice,AnthropicChoiceDict,GoogleChoiceDict,OpenAIChoiceDict,XAIChoiceDict]]]) – LLMs for prediction. Defaults to gen_llmc.config (
WeightTrainerConfig|dict|None) – Weight training configuration.gen_temperature (
float) – Sampling temperature for generation.predict_temperature (
float) – Sampling temperature for prediction.llm_semaphore_limit (
int) – Max concurrent LLM calls.max_policy_length (
int) – Max total policies to induce (< 500).class_ratio (
Tuple[float,float]) – Target YES/NO mix per generation batch, not the dataset’s actual ratio. Generation stops once either class can no longer fill its share, so imbalanced datasets won’t have every majority-class row shown during generation.max_samples_as_context (
int) – Samples per generation batch (max 100).p_predict_update_interval (
int) – Log progress every N policies during scoring.save_path (
str|PathLike[str] |None) – Directory for checkpoints and saved models.name (
str|None) – Instance name (alphanumeric + underscores only).random_state (
int) – Base random seed.confirm_requests (
bool) – Before fit()/predict() make any LLM calls, print an estimated request count per model and require a y/n confirmation.
- classmethod load(dir_path)#
Load a previously saved PolicyInduction instance.
- async fit(X, y)#
Fit the PolicyInduction model.
Runs policy generation, scoring, and weight fitting in sequence. A single checkpoint file spans generation and scoring, so resuming after an interruption never restores scores for a different set of policies than the ones actually in memory.
- get_memory()#
Return the policy memory DataFrame (policy text + predictions).
- Return type:
DataFrame
- async predict(samples)#
Yield predictions for each sample in the DataFrame.
Automatically checkpoints to self.save_path as it goes, resuming any matching in-progress checkpoint found there — same unconditional, single-file design as fit(). The checkpoint is deleted once every sample has been predicted.
- Parameters:
samples (
DataFrame) – DataFrame of samples to classify.- Yields:
(sample_index, policy_vector, prediction, token_counter)
- Return type:
AsyncGenerator[Tuple[Any,TypeAliasType,Literal['YES','NO'],TokenCounter],None]
- save(dir_path=None, for_production=False)#
Persist model state to disk.
Layout:
policy_induction.json manifest, config, state policies.parquet policy texts policy_predictions.parquet scored YES/NO matrix (dev only) data.parquet training data (dev only) lr.joblib trained logistic regression report.md human-readable fit summary
- async set_task(task_description, instructions_template=None)#
Set the task description and obtain the policy generation template.
Either accepts a custom template or generates one via LLM from the task description. The template must contain ‘<max_policy_length>’.
- property lr: LogisticRegression#
- property token_usage: TokenCounter#
- class WeightTrainerConfig(beta=0.5, penalty='l1', cv_folds=5, Cs=(0.001, 0.01, 0.1, 1, 10, 100, 1000), threshold_grid=(np.float64(0.01), np.float64(0.02), np.float64(0.03), np.float64(0.04), np.float64(0.05), np.float64(0.060000000000000005), np.float64(0.06999999999999999), np.float64(0.08), np.float64(0.09), np.float64(0.09999999999999999), np.float64(0.11), np.float64(0.12), np.float64(0.13), np.float64(0.14), np.float64(0.15000000000000002), np.float64(0.16), np.float64(0.17), np.float64(0.18000000000000002), np.float64(0.19), np.float64(0.2), np.float64(0.21000000000000002), np.float64(0.22), np.float64(0.23), np.float64(0.24000000000000002), np.float64(0.25), np.float64(0.26), np.float64(0.27), np.float64(0.28), np.float64(0.29000000000000004), np.float64(0.3), np.float64(0.31), np.float64(0.32), np.float64(0.33), np.float64(0.34), np.float64(0.35000000000000003), np.float64(0.36000000000000004), np.float64(0.37), np.float64(0.38), np.float64(0.39), np.float64(0.4), np.float64(0.41000000000000003), np.float64(0.42000000000000004), np.float64(0.43), np.float64(0.44), np.float64(0.45), np.float64(0.46), np.float64(0.47000000000000003), np.float64(0.48000000000000004), np.float64(0.49), np.float64(0.5), np.float64(0.51), np.float64(0.52), np.float64(0.53), np.float64(0.54), np.float64(0.55), np.float64(0.56), np.float64(0.5700000000000001), np.float64(0.5800000000000001), np.float64(0.59), np.float64(0.6), np.float64(0.61), np.float64(0.62), np.float64(0.63), np.float64(0.64), np.float64(0.65), np.float64(0.66), np.float64(0.67), np.float64(0.68), np.float64(0.6900000000000001), np.float64(0.7000000000000001), np.float64(0.7100000000000001), np.float64(0.72), np.float64(0.73), np.float64(0.74), np.float64(0.75), np.float64(0.76), np.float64(0.77), np.float64(0.78), np.float64(0.79), np.float64(0.8), np.float64(0.81), np.float64(0.8200000000000001), np.float64(0.8300000000000001), np.float64(0.8400000000000001), np.float64(0.85), np.float64(0.86), np.float64(0.87), np.float64(0.88), np.float64(0.89), np.float64(0.9), np.float64(0.91), np.float64(0.92), np.float64(0.93), np.float64(0.9400000000000001), np.float64(0.9500000000000001), np.float64(0.9600000000000001), np.float64(0.97), np.float64(0.98), np.float64(0.99)), class_weight_balanced=False, random_state=0)#
Bases:
objectConfiguration for training and optimizing ensemble weights.
- Parameters:
beta (
float) – Beta for F-beta score (e.g. 0.5 weights precision more).penalty (
Literal['l1','l2']) – Regularization type for logistic regression.cv_folds (
int) – Number of StratifiedKFold splits.threshold_grid (
Iterable[float]) – Decision thresholds to search over CV folds.class_weight_balanced (
bool) – Use class_weight=’balanced’ in LR.random_state (
int) – Random seed for CV splits.
-
threshold_grid:
Iterable[float] = (np.float64(0.01), np.float64(0.02), np.float64(0.03), np.float64(0.04), np.float64(0.05), np.float64(0.060000000000000005), np.float64(0.06999999999999999), np.float64(0.08), np.float64(0.09), np.float64(0.09999999999999999), np.float64(0.11), np.float64(0.12), np.float64(0.13), np.float64(0.14), np.float64(0.15000000000000002), np.float64(0.16), np.float64(0.17), np.float64(0.18000000000000002), np.float64(0.19), np.float64(0.2), np.float64(0.21000000000000002), np.float64(0.22), np.float64(0.23), np.float64(0.24000000000000002), np.float64(0.25), np.float64(0.26), np.float64(0.27), np.float64(0.28), np.float64(0.29000000000000004), np.float64(0.3), np.float64(0.31), np.float64(0.32), np.float64(0.33), np.float64(0.34), np.float64(0.35000000000000003), np.float64(0.36000000000000004), np.float64(0.37), np.float64(0.38), np.float64(0.39), np.float64(0.4), np.float64(0.41000000000000003), np.float64(0.42000000000000004), np.float64(0.43), np.float64(0.44), np.float64(0.45), np.float64(0.46), np.float64(0.47000000000000003), np.float64(0.48000000000000004), np.float64(0.49), np.float64(0.5), np.float64(0.51), np.float64(0.52), np.float64(0.53), np.float64(0.54), np.float64(0.55), np.float64(0.56), np.float64(0.5700000000000001), np.float64(0.5800000000000001), np.float64(0.59), np.float64(0.6), np.float64(0.61), np.float64(0.62), np.float64(0.63), np.float64(0.64), np.float64(0.65), np.float64(0.66), np.float64(0.67), np.float64(0.68), np.float64(0.6900000000000001), np.float64(0.7000000000000001), np.float64(0.7100000000000001), np.float64(0.72), np.float64(0.73), np.float64(0.74), np.float64(0.75), np.float64(0.76), np.float64(0.77), np.float64(0.78), np.float64(0.79), np.float64(0.8), np.float64(0.81), np.float64(0.8200000000000001), np.float64(0.8300000000000001), np.float64(0.8400000000000001), np.float64(0.85), np.float64(0.86), np.float64(0.87), np.float64(0.88), np.float64(0.89), np.float64(0.9), np.float64(0.91), np.float64(0.92), np.float64(0.93), np.float64(0.9400000000000001), np.float64(0.9500000000000001), np.float64(0.9600000000000001), np.float64(0.97), np.float64(0.98), np.float64(0.99))#