first commit
This commit is contained in:
@@ -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"]
|
||||
@@ -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()
|
||||
)
|
||||
@@ -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),
|
||||
}
|
||||
Reference in New Issue
Block a user