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
+4
View File
@@ -0,0 +1,4 @@
from app.experiments.matrix import generate_full_factorial_matrix
from app.experiments.runner import ExperimentRunner
__all__ = ["generate_full_factorial_matrix", "ExperimentRunner"]
+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()
)
+112
View File
@@ -0,0 +1,112 @@
"""Async experiment runner with retry, isolation, and resume support."""
import asyncio
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional
from sqlalchemy.orm import Session
from app.dataset.diff_extractor import extract_diff_bundle
from app.model_adapters.factory import create_adapter
from app.models import ExperimentRun, Result
from app.prompts.service import PromptService
class ExperimentRunner:
def __init__(
self,
db: Session,
concurrency: int = 5,
max_retries: int = 3,
strategy_id: str = "code_review",
):
self.db = db
self.concurrency = concurrency
self.max_retries = max_retries
self.strategy_id = strategy_id
self.prompt_service = PromptService(db)
def _get_pending_runs(self, experiment_id: Optional[str] = None) -> List[ExperimentRun]:
query = self.db.query(ExperimentRun)
if experiment_id:
query = query.filter_by(experiment_id=experiment_id)
return query.filter(ExperimentRun.status.in_(["pending", "failed"])).all()
def _reset_stale_running(self) -> None:
"""Mark runs stuck in running without result back to pending."""
stale = (
self.db.query(ExperimentRun)
.filter_by(status="running")
.filter(~ExperimentRun.result.has())
.all()
)
for run in stale:
run.status = "pending"
self.db.commit()
async def run_experiment(
self,
experiment_id: Optional[str] = None,
progress_callback=None,
) -> Dict[str, Any]:
self._reset_stale_running()
pending = self._get_pending_runs(experiment_id)
semaphore = asyncio.Semaphore(self.concurrency)
async def execute(run: ExperimentRun):
async with semaphore:
return await self._execute_run(run)
tasks = [asyncio.create_task(execute(run)) for run in pending]
results = []
for coro in asyncio.as_completed(tasks):
result = await coro
results.append(result)
if progress_callback:
progress_callback(result)
return {"total": len(pending), "completed": len(results)}
async def _execute_run(self, run: ExperimentRun) -> Dict[str, Any]:
run.status = "running"
run.started_at = datetime.now(timezone.utc)
self.db.commit()
try:
sample = run.sample
prompt = self.prompt_service.render(
self.strategy_id,
run.template_version.version_number,
{"language": sample.language, "diff": sample.diff},
)
adapter = create_adapter(run.model_id)
response = await adapter.chat(prompt, run.sampling_params)
run.status = "done"
run.completed_at = datetime.now(timezone.utc)
result = Result(
run_id=run.id,
raw_output=response.text,
token_usage=response.token_usage,
latency_ms=response.latency_ms,
)
self.db.add(result)
self.db.commit()
return {
"run_id": str(run.id),
"status": "done",
"model_id": run.model_id,
}
except Exception as e:
run.retry_count += 1
if run.retry_count > self.max_retries:
run.status = "failed"
else:
run.status = "pending" # will be retried on next resume
run.completed_at = datetime.now(timezone.utc)
self.db.commit()
return {
"run_id": str(run.id),
"status": run.status,
"error": str(e),
}