82 lines
2.5 KiB
Python
82 lines
2.5 KiB
Python
"""Full-factorial experiment matrix generation."""
|
||
|
||
from itertools import product
|
||
from typing import Any, Dict, List
|
||
from uuid import uuid4
|
||
|
||
from sqlalchemy.orm import Session
|
||
|
||
from app.models import Experiment, ExperimentRun, PromptTemplateVersion, Sample
|
||
from app.prompts.service import PromptService
|
||
|
||
|
||
def generate_full_factorial_matrix(
|
||
db: Session,
|
||
name: str,
|
||
models: List[str],
|
||
levels: List[str],
|
||
sample_ids: List[str],
|
||
repeats: int = 3,
|
||
sampling_params: Dict[str, Any] = None,
|
||
strategy_id: str = "code_review",
|
||
) -> Experiment:
|
||
"""Generate experiment runs for a full-factorial matrix.
|
||
|
||
Models × Levels × Samples × Repeats.
|
||
"""
|
||
sampling_params = sampling_params or {"temperature": 0.7, "max_tokens": 2048}
|
||
experiment = Experiment(
|
||
name=name,
|
||
models=models,
|
||
levels=levels,
|
||
sample_ids=sample_ids,
|
||
repeats=repeats,
|
||
sampling_params=sampling_params,
|
||
status="pending",
|
||
)
|
||
db.add(experiment)
|
||
db.flush()
|
||
|
||
prompt_service = PromptService(db)
|
||
runs = []
|
||
for model_id, level in product(models, levels):
|
||
# Resolve latest template version for this level
|
||
versions = prompt_service.list_versions(strategy_id, level)
|
||
if not versions:
|
||
raise ValueError(f"No prompt template found for {strategy_id}/{level}")
|
||
template_version = versions[-1]
|
||
|
||
for sample_id in sample_ids:
|
||
from uuid import UUID
|
||
|
||
sample_uuid = UUID(sample_id) if isinstance(sample_id, str) else sample_id
|
||
sample = db.query(Sample).filter_by(id=sample_uuid).first()
|
||
if not sample:
|
||
raise ValueError(f"Sample not found: {sample_id}")
|
||
|
||
for repeat_index in range(1, repeats + 1):
|
||
run = ExperimentRun(
|
||
run_id=str(uuid4()),
|
||
experiment_id=experiment.id,
|
||
sample_id=sample.id,
|
||
model_id=model_id,
|
||
template_version_id=template_version.id,
|
||
repeat_index=repeat_index,
|
||
status="pending",
|
||
sampling_params=sampling_params,
|
||
)
|
||
runs.append(run)
|
||
|
||
db.add_all(runs)
|
||
db.commit()
|
||
db.refresh(experiment)
|
||
return experiment
|
||
|
||
|
||
def count_pending_runs(db: Session, experiment_id: str) -> int:
|
||
return (
|
||
db.query(ExperimentRun)
|
||
.filter_by(experiment_id=experiment_id, status="pending")
|
||
.count()
|
||
)
|