Random Rule Forest (RRF)#
Module contents#
Random Rule Forest.
An interpretable ensemble framework for binary classification based on YES/NO questions generated by LLMs.
- class CVResult(fold_metrics, per_founder, summary)#
Bases:
objectResults from cross-validated aggregation evaluation.
- fold_metrics#
One row per (repeat, fold). Columns:
repeat,fold,k,t,precision,recall,f1,f_beta,accuracy,n_train,n_test.
- per_founder#
One row per (sample, repeat). Columns:
sample_idx,repeat,fold,y_true,y_pred,yes_count.
- summary#
Mean and standard deviation of each metric across folds. Keys follow the pattern
"<metric>_mean"and"<metric>_std"(e.g."f_beta_mean").
-
fold_metrics:
DataFrame#
-
per_founder:
DataFrame#
- class CostSensitiveConfig(screening_fraction=0.05, max_screening_samples=500, screening_metric='f1', screening_baseline='majority', max_questions_full_eval=20, enable_semantic_filter=True, semantic_threshold=0.85, semantic_emb_model='hashed_bag_of_words')#
Bases:
objectConfiguration for cost-sensitive training.
Cost-sensitive mode reduces LLM API costs through a multi-stage pipeline: 1. Auto semantic filtering (optional) 2. Screening evaluation on small subset 3. Pruning low-performing questions 4. Top-N selection 5. Full evaluation on complete training set
For large datasets (9000+ samples), this can reduce costs by 44-74%.
- Parameters:
-
max_questions_full_eval:
int= 20# Maximum number of top questions to evaluate on full training set.
After screening, only the top N questions by screening_metric are fully evaluated. Lower values = more aggressive cost reduction.
-
max_screening_samples:
int|None= 500# Maximum number of samples for screening, regardless of fraction.
-
screening_baseline:
float|str= 'majority'# Baseline for pruning questions.
“majority”: F1 score of majority-class classifier
“random”: Expected F1 of random classifier
float: Explicit threshold (e.g., 0.6)
Questions scoring at or below baseline are excluded.
-
screening_fraction:
float= 0.05# 5%).
- Type:
Fraction of training data to use for screening (default
-
screening_metric:
str= 'f1'# Metric to use for screening evaluation (“f1”, “precision”, or “recall”).
- class PromptPreset(name, description, question_gen_system, question_gen_user_template, question_answer_system, question_answer_user_template)#
Bases:
objectNamed prompt collection for a specific RRF domain.
A preset provides domain-specific system messages and user templates that bypass the default meta-prompt step, giving tighter control over question generation and answering behaviour.
- Parameters:
- name#
Short identifier (used for registry lookup).
- description#
Human-readable description of the preset.
- question_gen_system#
System message for the question-generation LLM call.
- question_gen_user_template#
User prompt template for question generation. Must contain
{num_questions}and{samples}placeholders.
- question_answer_system#
System message for the question-answering LLM call.
- question_answer_user_template#
User prompt template for question answering. Must contain
{question}and{sample}placeholders.
- class QuestionExclusion(*values)#
Bases:
StrEnum- COST_PRUNING = 'cost_pruning'#
- EXPERT = 'expert'#
- PREDICTION_SIMILARITY = 'prediction_similarity'#
- SEMANTICS = 'semantics'#
- class RRF(qgen_llmc, qanswer_llmc=None, qgen_temperature=0.0, qanswer_temperature=0.0, llm_semaphore_limit=3, answer_similarity_func='hamming', max_generated_questions=100, max_samples_as_context=30, class_ratio=(1.0, 1.0), q_answer_update_interval=10, save_path=None, name=None, random_state=42, use_cumulative_memory=True, qanswer_batch_size=None, question_scoring_f_beta=1.0, semantic_filtering_during_fit=False, semantic_similarity_threshold=0.9, aggregation_metric='f1', aggregation_max_k=None, aggregation_method='vote', elasticnet_cs=(0.05, 0.1, 0.5), elasticnet_l1_ratios=(0.1, 0.5), elasticnet_cv=3, cost_sensitive=False, cost_sensitive_config=None, prompt_preset=None, _llm=None)#
Bases:
objectInterpretable ensemble binary classifier.
- Parameters:
qgen_llmc (
List[Union[AnthropicChoice,GoogleChoice,OpenAIChoice,XAIChoice,AnthropicChoiceDict,GoogleChoiceDict,OpenAIChoiceDict,XAIChoiceDict]]) – LLMs to use for question generation, in priority order.qanswer_llmc (
Optional[List[Union[AnthropicChoice,GoogleChoice,OpenAIChoice,XAIChoice,AnthropicChoiceDict,GoogleChoiceDict,OpenAIChoiceDict,XAIChoiceDict]]]) – LLMs to use for answering questions, in priority order. If None, use qgen_llmc.qgen_temperature (
float) – Sampling temperature for question generation.qanswer_temperature (
float) – Sampling temperature for answering questions.llm_semaphore_limit (
int) – Max concurrent LLM calls.answer_similarity_func (
Union[str,Literal['jaccard','hamming','correlation']]) – Function to use for answer similarity.max_generated_questions (
int) – Maximum number of questions to generate. Max 1000max_samples_as_context (
int) – Number of samples used as context in a round of question generation. max 100 Max 100class_ratio (
Tuple[float,float]) – Ratio of YES to NO samples to use as context in a round of question generation.q_answer_update_interval (
int) – Logging interval of question answering.save_path (
str|PathLike[str] |None) – Directory to save checkpoints/models.random_state (
int) – Random seed.use_cumulative_memory (
bool) – Whether to use cumulative memory when generating questions across multiple LLM calls.qanswer_batch_size (
int|None) – Maximum number of samples to answer in a single LLM call. If None or 1, batching is disabled and the original per-sample behaviour is used (one LLM call per sample). Set >1 to enable true batched answering.question_scoring_f_beta (
float) – Beta parameter for computing F-beta score on questions. Default 1.0 (F1). Use 0.5 to weight precision more, or 2.0 to weight recall more. Must be > 0.semantic_filtering_during_fit (
bool) – If True, run semantic deduplication on generated questions before the expensive answering step. Useshashed_bag_of_wordsembeddings (no API calls). Default False.semantic_similarity_threshold (
float) – Cosine-similarity threshold used for early semantic filtering. Only relevant whensemantic_filtering_during_fit=True. Default 0.9.aggregation_metric (
Literal['f1','f_beta','accuracy','precision','recall']) – Metric optimized when tuning (K, T) for founder-level prediction duringfit(). One of"f1","f_beta","accuracy","precision","recall". When"f_beta"is selected, usesquestion_scoring_f_betaas the beta parameter. Default"f1".aggregation_max_k (
int|None) – Maximum K (number of top questions) to consider during (K, T) grid search.Nonemeans use all active questions. DefaultNone(vote method only).aggregation_method (
Literal['vote','elasticnet']) – How per-question answers are combined into a founder-level label."vote"(default) uses the unit-weight top-(K, T) scheme tuned by grid search."elasticnet"instead fits an elastic-net logistic regression over all active questions, learning a signed weight per question plus a decision threshold. Learned weights are usually more accurate (they down-weight noisy questions and exploit correlations) but trade away the simple “N of K rules fired” interpretation for inspectable logistic coefficients.aggregation_max_kand thek/toverrides onpredict_founder_levelapply to"vote"only.elasticnet_cs (
Tuple[float,...]) – Inverse-regularisation grid (sklearnCs) searched by inner CV whenaggregation_method="elasticnet". Default(0.05, 0.1, 0.5).elasticnet_l1_ratios (
Tuple[float,...]) – Elastic-net mixing grid (0 = pure L2, 1 = pure L1) searched by inner CV. Default(0.1, 0.5).elasticnet_cv (
int) – Number of inner CV folds for elastic-net hyperparameter selection. Must be >= 2. Default3.cost_sensitive (
bool) – Enable cost-sensitive mode with screening and early pruning.cost_sensitive_config (
CostSensitiveConfig|None) – Configuration for cost-sensitive mode. If None, uses default CostSensitiveConfig.prompt_preset (
str|PromptPreset|None) – Optional prompt preset (PromptPresetinstance or registered name string). When provided, bypasses the meta-prompt step and uses domain-specific prompts for generation and answering._llm (
Any) – LLM instance for testing (dependency injection). If None, uses global llm.
- classmethod load(dir_path)#
Load an RRF saved by save.
- static aggregate_predictions(response_matrix, question_scores, k, t)#
Aggregate per-question binary responses into founder-level labels.
Selects the top-K questions by score (descending) and predicts YES for a founder if at least T of those K questions are answered YES.
- Parameters:
- Return type:
Series- Returns:
Series of “YES”/”NO” predictions indexed by sample.
- Raises:
ValueError – If k or t are out of valid range.
- async add_question(question)#
Add a question to the RRF.
- Parameters:
question (
str) – The question to add.- Raises:
ValueError – If question already exists.
- Return type:
Literal[True]
- exclusion_report(as_dict=False)#
Return a structured summary of excluded questions and why they were dropped.
- filter_questions_on_pred_similarity(threshold)#
Filter questions on prediction similarity.
If two questions have a prediction similarity greater than or equal to the threshold, the question with the lower f1 score is excluded.
- Parameters:
- Raises:
AssertionError – If threshold is not between > 0 and <= 1.
- Return type:
- async filter_questions_on_semantics(threshold, emb_model)#
Filter questions on semantics.
If two questions have a semantic similarity greater than or equal to the threshold, the question with the lower f1 score is excluded.
- Parameters:
- Raises:
AssertionError – If threshold is not between > 0 and <= 1.
- Return type:
- async fit(X=None, y=None, *, X_val=None, y_val=None, copy_data=True, reset=False)#
Fit the RRF to the data.
Generates questions, answers them on the provided data, computes per-question metrics, and tunes (K, T) for founder-level aggregation via
predict_founder_level().Note
(K, T) are tuned on the data passed to
fit(). For unbiased evaluation, fit on training data and evaluate on a held-out test set.- Parameters:
X (
DataFrame|None) – Training features. Required on first run or with reset=True.y (
Optional[Sequence[str]]) – Training labels. Required on first run or with reset=True.X_val (
DataFrame|None) – Optional validation features for cost-sensitive mode.y_val (
Optional[Sequence[str]]) – Optional validation labels for cost-sensitive mode.copy_data (
bool) – Whether to copy input data.reset (
bool) – Clear existing state and restart forest generation.
- Returns:
Updated RRF.
- Return type:
Self
- Raises:
ValueError – If data requirements aren’t met or invalid reset usage.
- get_answers()#
Get answers dataframe.
- Returns:
Columns as the questions ids (not excluded during semantic filtering).
Index as the samples indices from self._X.
- Return type:
DataFrame containing
- get_questions()#
Get generated questions dataframe.
- Returns:
question: Generated question text
embedding: Question embedding vectors
exclusion: Question exclusion principle
precision: Precision for each question
recall: Recall for each question
f1_score: F1 score for each question
- f_beta_score: F-beta score for each question (configurable via
question_scoring_f_beta; defaults to F1 when beta=1.0)
accuracy: Accuracy for each question
- Return type:
DataFrame containing
- async predict(samples, *, max_concurrent=None, checkpoint_path=None, checkpoint_every=None, resume=False)#
Predict labels for samples.
Uses batched LLM answering to reduce API calls. Each batch groups multiple samples into a single LLM call per question. The batch size is controlled by
qanswer_batch_size(default 20 when not set).- Parameters:
samples (
DataFrame) – Samples to predict.max_concurrent (
int|None) – Maximum number of questions to process concurrently.None(default) uses a sequential loop.checkpoint_path (
str|PathLike[str] |None) – Directory wherepredict_checkpoint.jsonis written.Nonedisables checkpointing.checkpoint_every (
int|None) – Save a checkpoint every N completed questions. Requires checkpoint_path to be set.resume (
bool) – IfTrue, load an existing checkpoint from checkpoint_path and skip already-completed questions.
- Return type:
AsyncGenerator[Tuple[Any,str,Literal['YES','NO'],TokenCounter],None]- Returns:
Generator of predictions[sample_index, question, answer, token_counter]
- Raises:
ValueError – If samples is empty or does not have the correct column.
- async predict_founder_level(X, *, k=None, t=None)#
Founder-level binary predictions using top-K / threshold-T.
Calls
predict()internally to get per-question answers, then aggregates using the top-K questions (ranked by f_beta_score) and a YES-count threshold T.K and T are learned during
fit()via grid search on training data. You can override them with explicit arguments.Note
K/T are tuned on training data (same data used for question scoring). For stricter separation, pass explicit
kandtvalues tuned on a held-out validation set.Example:
# sklearn-style workflow rrf = RRF(qgen_llmc=llm_choices, name="my_rrf") await rrf.set_tasks(task_description="Classify founders") await rrf.fit(X_train, y_train) # Predict on new data results = await rrf.predict_founder_level(X_test) print(results[["prediction", "yes_count"]])
- Parameters:
- Return type:
DataFrame- Returns:
DataFrame indexed by
X.index. Foraggregation_method="vote"(default) the columns areprediction(“YES”/”NO”),yes_count(YES answers among the top-K questions),kandt. Foraggregation_method="elasticnet"the columns areprediction,probability(learned P(YES)) andthreshold; thek/targuments are ignored in that mode.- Raises:
ValueError – If (vote mode) k/t are not provided and not learned during fit, or (elasticnet mode) the model is not fitted.
- save(dir_path=None, for_production=False)#
Save model config to JSON and dataframes to parquet in a directory.
If dir_path is None, uses <self.save_path>/<self.name>. If for_production is True, strips the questions dataframe and does not save the answers and training dataframes.
- async set_tasks(instructions_template=None, task_description=None)#
Initialize question generation instructions template.
This sets the task description for the RRF. Either sets a custom template or generates one from task description using LLM. For most users, LLM generation is recommended over custom templates.
- Parameters:
- Return type:
- Returns:
The question generation instructions template.
- Raises:
ValueError – If template missing required tag or generation fails.
AssertionError – If both parameters are None.
- async update_question_exclusion(question_id, exclusion)#
Update a question exclusion.
- Parameters:
question_id (
str) – The id of the question to update.exclusion (
QuestionExclusion|None) – The exclusion to set. If None, removes the exclusion status.
- Returns:
The question that was updated.
- Return type:
- Raises:
ValueError – If question id is not found.
- property question_gen_instructions_template: str | None#
Get the question generation instructions template.
- property token_usage: TokenCounter#
Get the token counter for the RRF.
- cross_validate_aggregation(answer_matrix, y, *, n_splits=10, n_repeats=10, metric='f_beta', beta=0.5, max_k=None, random_state=42)#
Evaluate RRF aggregation via repeated stratified k-fold CV.
For each fold the function:
Scores every question on the training split (f-beta).
Ranks questions by that score (descending).
Grid-searches (K, T) on the training split.
Evaluates with the chosen (K, T) on the test split.
No LLM calls are made — everything is computed from the pre-built answer matrix.
- Parameters:
answer_matrix (
DataFrame) –(n_samples, n_questions)DataFrame with"YES"/"NO"values. Must not include samples whose labels were seen during question generation.y (
Sequence[str]) – True labels ("YES"/"NO"), one per sample.n_splits (
int) – Number of folds per repeat.n_repeats (
int) – Number of times to repeat the k-fold split.metric (
str) – Metric to optimise when tuning (K, T). Any value accepted byRRF._compute_metric()(e.g."f_beta","f1","precision").beta (
float) – Beta parameter for F-beta scoring and tuning.max_k (
int|None) – Upper bound on K during grid search.None= use all questions.random_state (
int) – Base random seed; each repeat usesrandom_state + repeat.
- Return type:
- Returns:
A
CVResultwith fold-level metrics, per-founder predictions, and an aggregated summary.