Files
2026-09-19 12:54:45 +08:00

82 lines
2.5 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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()
)