first commit

This commit is contained in:
eeymoo
2026-09-19 12:54:45 +08:00
commit 6fc5b64077
126 changed files with 8601 additions and 0 deletions
+81
View File
@@ -0,0 +1,81 @@
"""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()
)