OpenEEGBench / backend /app /services /submissions.py
bruAristimunha's picture
feat: OpenEEGBench submission flow (arena runs the eval)
4db3be0
Raw
History Blame Contribute Delete
5.49 kB
"""Pure helpers to validate and build OpenEEGBench eval submissions."""
from __future__ import annotations
import re
from datetime import datetime, timezone
from typing import Any, Dict, List
from app.config.submission_config import (
FINETUNING_STRATEGIES, PEFT_STRATEGIES, IA3_STRATEGY, HEADS,
DATASET_KEYS, DEFAULT_STRATEGIES, DEFAULT_HEADS, DEFAULT_N_SEEDS,
DEFAULT_HEAD_MODULE_NAME, MONTHLY_QUOTA, FRAMEWORK,
)
_DOTTED_RE = re.compile(r"^[A-Za-z_]\w*(\.[A-Za-z_]\w*)+$")
_EMAIL_RE = re.compile(r"^[^@\s]+@[^@\s]+\.[^@\s]+$")
def _as_list(value) -> list:
if value is None:
return []
return value if isinstance(value, list) else [value]
def _str(payload, key) -> str:
return (payload.get(key) or "").strip()
def validate_submission(payload: Dict[str, Any]) -> List[str]:
"""Return human-readable validation errors (empty list == valid)."""
errors: List[str] = []
model_cls = _str(payload, "model_cls")
if not model_cls:
errors.append("model_cls is required (dotted path, e.g. braindecode.models.BIOT).")
elif not _DOTTED_RE.match(model_cls):
errors.append(f"model_cls '{model_cls}' is not a valid dotted import path.")
if not _str(payload, "model_name"):
errors.append("model_name is required.")
if not _str(payload, "organization"):
errors.append("organization is required.")
hub_repo = _str(payload, "hub_repo")
checkpoint_url = _str(payload, "checkpoint_url")
if not hub_repo and not checkpoint_url:
errors.append("Provide a weights source: either hub_repo or checkpoint_url.")
if hub_repo and checkpoint_url:
errors.append("Provide only one weights source (hub_repo or checkpoint_url).")
strategies = _as_list(payload.get("finetuning_strategies")) or DEFAULT_STRATEGIES
bad = [s for s in strategies if s not in FINETUNING_STRATEGIES]
if bad:
errors.append(f"Unknown finetuning strategies: {bad}. Valid: {FINETUNING_STRATEGIES}.")
heads = _as_list(payload.get("heads")) or DEFAULT_HEADS
bad = [h for h in heads if h not in HEADS]
if bad:
errors.append(f"Unknown heads: {bad}. Valid: {HEADS}.")
datasets = _as_list(payload.get("datasets")) or DATASET_KEYS
bad = [d for d in datasets if d not in DATASET_KEYS]
if bad:
errors.append(f"Unknown datasets: {bad}. Valid: {DATASET_KEYS}.")
if any(s in PEFT_STRATEGIES for s in strategies) and not _as_list(payload.get("peft_target_modules")):
errors.append("peft_target_modules is required when a PEFT strategy (lora/ia3/adalora/dora/oft) is selected.")
if IA3_STRATEGY in strategies and not _as_list(payload.get("peft_ff_modules")):
errors.append("peft_ff_modules is required when the ia3 strategy is selected.")
n_seeds = payload.get("n_seeds", DEFAULT_N_SEEDS)
if not isinstance(n_seeds, int) or isinstance(n_seeds, bool) or n_seeds < 1:
errors.append("n_seeds must be an integer >= 1.")
email = _str(payload, "contact_email")
if email and not _EMAIL_RE.match(email):
errors.append("contact_email is not a valid email address.")
if payload.get("visibility", "Public") not in ("Public", "Private"):
errors.append("visibility must be 'Public' or 'Private'.")
return errors
def build_request_entry(payload: Dict[str, Any], submitter: str, now: datetime | None = None) -> Dict[str, Any]:
now = now or datetime.now(timezone.utc)
strategies = _as_list(payload.get("finetuning_strategies")) or DEFAULT_STRATEGIES
heads = _as_list(payload.get("heads")) or DEFAULT_HEADS
datasets = _as_list(payload.get("datasets")) or DATASET_KEYS
return {
"schema_version": 1,
"framework": FRAMEWORK,
"model_name": _str(payload, "model_name"),
"organization": _str(payload, "organization"),
"submitter": submitter,
"contact_email": _str(payload, "contact_email"),
"model_url": _str(payload, "model_url"),
"paper_url": _str(payload, "paper_url"),
"additional_info": _str(payload, "additional_info"),
"visibility": payload.get("visibility", "Public"),
"benchmark_kwargs": {
"model_cls": _str(payload, "model_cls"),
"hub_repo": _str(payload, "hub_repo") or None,
"checkpoint_url": _str(payload, "checkpoint_url") or None,
"model_kwargs": payload.get("model_kwargs") or {},
"head_module_name": payload.get("head_module_name") or DEFAULT_HEAD_MODULE_NAME,
"peft_target_modules": _as_list(payload.get("peft_target_modules")),
"peft_ff_modules": _as_list(payload.get("peft_ff_modules")),
"finetuning_strategies": strategies,
"heads": heads,
"datasets": datasets,
"n_seeds": int(payload.get("n_seeds", DEFAULT_N_SEEDS)),
},
"status": "PENDING",
"submitted_time": now.strftime("%Y-%m-%dT%H:%M:%SZ"),
"job_id": -1,
}
def count_user_submissions_this_month(entries, submitter: str, now: datetime | None = None) -> int:
now = now or datetime.now(timezone.utc)
month = now.strftime("%Y-%m")
return sum(
1 for e in entries
if e.get("submitter") == submitter and str(e.get("submitted_time", "")).startswith(month)
)
def monthly_quota_exceeded(entries, submitter: str, now: datetime | None = None) -> bool:
return count_user_submissions_this_month(entries, submitter, now) >= MONTHLY_QUOTA