Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
98 changes: 95 additions & 3 deletions distil/eval/results.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,10 +33,79 @@

logger = logging.getLogger("distil.eval.results")

MAX_REEVAL_FAILURES = 3


# ── Helpers ──────────────────────────────────────────────────────────────


def _merge_reeval_composite(
stored: dict,
fresh: dict,
stored_prompts: int,
fresh_prompts: int,
) -> dict:
"""Prompt-count-weighted merge of a re-eval composite into the stored one.

Per-axis scores are averaged weighted by prompt counts. Aggregates
(worst, worst_3_mean, weighted, final) are recomputed from the merged
axes using the same formula as ``compute_composite``.

If either side has no prompt count, fall back to equal weighting (0.5).
"""
total = stored_prompts + fresh_prompts
if total <= 0:
w_stored, w_fresh = 0.5, 0.5
else:
w_stored = stored_prompts / total
w_fresh = fresh_prompts / total

stored_axes = stored.get("axes") or {}
fresh_axes = fresh.get("axes") or {}
all_keys = set(stored_axes.keys()) | set(fresh_axes.keys())
merged_axes: dict[str, float | None] = {}
for k in all_keys:
sv = stored_axes.get(k)
fv = fresh_axes.get(k)
if sv is not None and fv is not None:
merged_axes[k] = round(w_stored * sv + w_fresh * fv, 4)
elif fv is not None:
merged_axes[k] = fv
else:
merged_axes[k] = sv

weights = {k: w for k, w in settings.axis_weights().items() if w > 0}
broken = set(fresh.get("broken_axes") or [])
ranked = {k: v for k, v in merged_axes.items() if v is not None and k in weights and k not in broken}
weighted_axes = {k: v for k, v in merged_axes.items() if v is not None and k in weights}

if not ranked:
return fresh

worst = min(ranked.values())
sorted_vals = sorted(ranked.values())
k_eff = min(settings.worst_3_mean_k, len(sorted_vals))
worst_k_mean = sum(sorted_vals[:k_eff]) / k_eff
total_w = sum(weights[k] for k in weighted_axes)
weighted_score = sum(weights[k] * v for k, v in weighted_axes.items()) / total_w if total_w else None
alpha = settings.composite_final_bottom_weight
if weighted_score is not None:
final = alpha * worst_k_mean + (1.0 - alpha) * weighted_score
else:
final = worst_k_mean

merged = dict(fresh)
merged["axes"] = merged_axes
merged["worst"] = round(worst, 4)
merged["worst_3_mean"] = round(worst_k_mean, 4)
merged["final"] = round(final, 4)
merged["final_alpha"] = round(alpha, 4)
merged["weighted"] = round(weighted_score, 4) if weighted_score is not None else None
merged["present_count"] = len(ranked)
merged["prompts_accumulated"] = total
return merged


def _resolve_anchor(rows: dict[str, dict], king_name: str | None, key: str) -> float | None:
"""Pick the anchor for relative axes (king if seated, else round-min).

Expand Down Expand Up @@ -328,6 +397,7 @@ def process_round(
coldkey = idx.get("coldkey")
commit_block = idx.get("commit_block")
is_king = (name == king_name)
is_reeval = bool(idx.get("is_reeval"))

# SHA256 exact-weight duplicate check (always-on, cross-round).
sha = row.get("weights_sha256")
Expand Down Expand Up @@ -431,7 +501,19 @@ def process_round(
f"dq={comp.get('disqualified')}) — preserving prior record"
)
else:
state.composite_scores[str(uid) if uid is not None else name] = comp
comp_key = str(uid) if uid is not None else name
if is_reeval:
stored = state.composite_scores.get(comp_key)
if stored and stored.get("final") is not None:
stored_prompts = int(stored.get("prompts_accumulated") or settings.eval_n_prompts)
fresh_prompts = int(row.get("n_prompts") or settings.eval_n_prompts)
comp = _merge_reeval_composite(stored, comp, stored_prompts, fresh_prompts)
logger.info(
f"uid={uid} ({name}): re-eval composite merged "
f"(stored_n={stored_prompts}, fresh_n={fresh_prompts}, "
f"final={comp.get('final')})"
)
state.composite_scores[comp_key] = comp

# Update the cross-round fingerprint store so future rounds see this uid.
if isinstance(fp, dict) and fp.get("layer_fingerprints"):
Expand Down Expand Up @@ -506,10 +588,11 @@ def process_round(
"kl": kl_val if kl_val != float("inf") else None,
"is_king": is_king,
"is_reference": is_ref,
"is_reeval": is_reeval,
"prompts_scored": prompts_scored,
"prompts_total": prompts_scored,
"paired_prompts": prompts_scored,
"dethrone_eligible": (not is_king) and (comp or {}).get("disqualified") is not True,
"dethrone_eligible": (not is_king) and (not is_reeval) and (comp or {}).get("disqualified") is not True,
"early_stopped": False,
"composite": comp,
"axes_summary": (comp or {}).get("axes"),
Expand Down Expand Up @@ -551,7 +634,9 @@ def process_round(
if uid is not None and not is_ref:
if load_succeeded:
state.reset_failures(int(uid))
elif load_failed:
if is_reeval:
state.reeval_failures.pop(str(uid), None)
elif load_failed and not is_reeval:
load_failures_after = state.record_failure(int(uid), name)
err_short = str(row.get("error") or "")[:200]
result_row["error"] = err_short
Expand Down Expand Up @@ -597,6 +682,13 @@ def process_round(
f"{load_failures_after} consecutive load failures "
f"on {name}: {err_short[:120]}"
)
elif load_failed and is_reeval:
strikes = state.reeval_failures.get(str(uid), 0) + 1
state.reeval_failures[str(uid)] = strikes
logger.warning(
f"uid={uid} ({name}): re-eval load failure "
f"(strike {strikes}/{MAX_REEVAL_FAILURES})"
)
if (comp or {}).get("disqualified"):
result_row["disqualified"] = True
result_row["dq_reason"] = (comp or {}).get("dq_reason")
Expand Down
107 changes: 106 additions & 1 deletion distil/eval/round.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,9 @@
or 10
)

MAX_PER_COLDKEY = 2
MAX_REEVAL_FAILURES = 3


def _model_key(c: Commitment) -> str:
return c.key
Expand Down Expand Up @@ -438,7 +441,6 @@ def select_challengers(
# from any one coldkey in any single round. Doesn't block
# registration (that's bittensor's gate) — only stops one
# coldkey from monopolizing scoring bandwidth.
MAX_PER_COLDKEY = 2
coldkey_counts: dict[str, int] = {}
if king_c is not None and getattr(king_c, "coldkey", None):
coldkey_counts[king_c.coldkey] = coldkey_counts.get(king_c.coldkey, 0) + 1
Expand Down Expand Up @@ -540,6 +542,96 @@ def _coldkey_cap_blocks(c: "Commitment") -> bool:
return accepted


def select_reevals(
commitments: dict[int, Commitment],
state: ValidatorState,
*,
king_uid: int | None,
n: int,
fresh_challengers: list[Commitment],
) -> list[Commitment]:
"""Backfill idle slots with top-N already-evaluated UIDs for re-eval.

Selection is by ``composite_scores[uid]["final"]`` descending. Only
UIDs with a valid on-chain commitment and an existing composite are
eligible. Re-eval UIDs are tagged ``is_reeval: true`` downstream so
they cannot trigger the dethrone gate.

The per-coldkey cap and ``model@revision`` dedup are shared with the
fresh-challenger set passed in via ``fresh_challengers``.
"""
if n <= 0:
return []

in_round_uids: set[int] = set()
if king_uid is not None:
in_round_uids.add(int(king_uid))
for c in fresh_challengers:
in_round_uids.add(int(c.uid))

seen_keys: set[str] = set()
coldkey_counts: dict[str, int] = {}
king_c = commitments.get(int(king_uid)) if king_uid is not None else None
if king_c is not None:
if getattr(king_c, "key", None):
seen_keys.add(king_c.key)
ck = getattr(king_c, "coldkey", None)
if ck:
coldkey_counts[ck] = coldkey_counts.get(ck, 0) + 1
for c in fresh_challengers:
key = getattr(c, "key", None)
if key:
seen_keys.add(key)
ck = getattr(c, "coldkey", None)
if ck:
coldkey_counts[ck] = coldkey_counts.get(ck, 0) + 1

composite_scores = state.composite_scores or {}
scored_uids: list[tuple[float, int]] = []
for uid_str, comp in composite_scores.items():
if not isinstance(comp, dict):
continue
final = comp.get("final")
if final is None:
continue
try:
scored_uids.append((float(final), int(uid_str)))
except (TypeError, ValueError):
continue
scored_uids.sort(reverse=True)

accepted: list[Commitment] = []
for _final, uid in scored_uids:
if len(accepted) >= n:
break
if uid in in_round_uids:
continue
if uid not in commitments:
continue
c = commitments[uid]
if state.is_disqualified(c.hotkey, uid=c.uid):
continue
if state.reeval_failures.get(str(uid), 0) >= MAX_REEVAL_FAILURES:
continue
key = getattr(c, "key", None)
if key and key in seen_keys:
continue
ck = getattr(c, "coldkey", None)
if ck and coldkey_counts.get(ck, 0) >= MAX_PER_COLDKEY:
continue
if key:
seen_keys.add(key)
if ck:
coldkey_counts[ck] = coldkey_counts.get(ck, 0) + 1
accepted.append(c)
if accepted:
logger.info(
f"select_reevals: backfilling {len(accepted)} re-eval slot(s): "
f"{[(int(c.uid), c.model) for c in accepted]}"
)
return accepted


def build_round_spec(
*,
block: int,
Expand All @@ -548,6 +640,7 @@ def build_round_spec(
reference_repo: str,
king: Commitment | None,
challengers: list[Commitment],
reevals: list[Commitment] | None = None,
) -> dict[str, Any]:
"""JSON-serializable spec uploaded to the GPU pod."""
students: list[dict[str, Any]] = []
Expand All @@ -573,6 +666,18 @@ def build_round_spec(
"is_king": False,
}
)
for c in (reevals or []):
students.append(
{
"name": _model_key(c),
"repo": c.model,
"revision": c.revision,
"uid": c.uid,
"hotkey": c.hotkey,
"is_king": False,
"is_reeval": True,
}
)
return {
"round_id": int(time.time()),
"block": block,
Expand Down
Loading