Spaces:
Running
Running
| import os | |
| import json | |
| import random | |
| import math | |
| import logging | |
| import traceback | |
| from pathlib import Path | |
| from dataclasses import dataclass | |
| from typing import Dict, List, Tuple, Optional | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import AutoTokenizer, AutoModelForCausalLM | |
| import gradio as gr | |
| import pandas as pd | |
| # --------------------------------------------------------------------------- | |
| # Logging | |
| # --------------------------------------------------------------------------- | |
| logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s") | |
| logger = logging.getLogger(__name__) | |
| # --------------------------------------------------------------------------- | |
| # Constants & Paths | |
| # --------------------------------------------------------------------------- | |
| MODEL_IDS: List[str] = [ | |
| "CodeSoft/MetaDiffusion-150M-ChatBase", | |
| "BananaMind/BananaMind-2-Medium-Chat", | |
| "SupraLabs/Supra2-100M-Instruct", | |
| "HuggingFaceTB/SmolLM2-135M-Instruct", | |
| ] | |
| MODEL_DISPLAY: Dict[str, str] = { | |
| "CodeSoft/MetaDiffusion-150M-ChatBase": "MetaDiffusion-150M-ChatBase", | |
| "BananaMind/BananaMind-2-Medium-Chat": "BananaMind-2-Medium-Chat", | |
| "SupraLabs/Supra2-100M-Instruct": "Supra2-100M-Instruct", | |
| "HuggingFaceTB/SmolLM2-135M-Instruct": "SmolLM2-135M-Instruct", | |
| } | |
| FALLBACK_IDS: Dict[str, str] = {} | |
| INIT_RATING = 1000 | |
| K_FACTOR = 32 | |
| SCALE = 400 | |
| BASE = 10 | |
| # All data in ./data | |
| try: | |
| BASE_DIR = Path(__file__).parent | |
| except NameError: | |
| BASE_DIR = Path(".") | |
| # Prefer /data (HF Space bucket mount) if available, otherwise fallback to ./data | |
| # Bucket is mounted at /data in Space — use dynamic check each call so late mounts are detected | |
| def get_data_dir() -> Path: | |
| bucket = Path("/data") | |
| if bucket.exists() and bucket.is_dir(): | |
| try: | |
| # Ensure writable (touch test) | |
| (bucket / ".write_test").touch(exist_ok=True) | |
| (bucket / ".write_test").unlink(missing_ok=True) | |
| return bucket | |
| except Exception: | |
| pass | |
| # Fallback to local ./data | |
| local = BASE_DIR / "data" | |
| try: | |
| local.mkdir(parents=True, exist_ok=True) | |
| except Exception: | |
| pass | |
| return local | |
| def get_elo_file() -> Path: | |
| return get_data_dir() / "elo.json" | |
| def get_chat_file() -> Path: | |
| return get_data_dir() / "chats.jsonl" | |
| # Keep legacy globals for backwards compat (now dynamic via functions) | |
| DATA_DIR = get_data_dir() | |
| ELO_FILE = get_elo_file() | |
| CHAT_FILE = get_chat_file() | |
| GEN_DEFAULTS: Dict[str, dict] = { | |
| "HuggingFaceTB/SmolLM2-135M-Instruct": {"max_new_tokens": 64, "temperature": 0.7, "top_p": 0.9, "repetition_penalty": 1.1, "do_sample": True}, | |
| "SupraLabs/Supra2-100M-Instruct": {"max_new_tokens": 64, "temperature": 0.7, "top_p": 0.9, "top_k": 25, "repetition_penalty": 1.1, "do_sample": True, "no_repeat_ngram_size": 3}, | |
| "BananaMind/BananaMind-2-Medium-Chat": {"max_new_tokens": 64, "temperature": 0.7, "top_p": 0.9, "repetition_penalty": 1.1, "do_sample": True}, | |
| "CodeSoft/MetaDiffusion-150M-ChatBase": {"max_new_tokens": 96, "num_steps": 128, "temperature": 0.7, "top_p": 0.9, "repetition_penalty": 1.5}, | |
| } | |
| MODEL_CONTEXT: Dict[str, int] = { | |
| "HuggingFaceTB/SmolLM2-135M-Instruct": 2048, | |
| "SupraLabs/Supra2-100M-Instruct": 1024, | |
| "BananaMind/BananaMind-2-Medium-Chat": 3072, | |
| "CodeSoft/MetaDiffusion-150M-ChatBase": 5120, | |
| } | |
| DEVICE = "cpu" | |
| class MetaDiffusionConfig: | |
| hidden_size: int = 768 | |
| intermediate_size: int = 2112 | |
| num_hidden_layers: int = 16 | |
| num_attention_heads: int = 12 | |
| num_key_value_heads: int = 6 | |
| head_dim: int = 64 | |
| vocab_size: int = 32000 | |
| mask_vocab_size: int = 32010 | |
| max_position_embeddings: int = 5120 | |
| rope_theta: float = 10000.0 | |
| rms_norm_eps: float = 1e-6 | |
| hidden_act: str = "silu" | |
| timestep_emb_hidden: int = 768 | |
| mask_token_id: int = 32000 | |
| pad_token_id: int = 1 | |
| mask_ratio_min: float = 0.0 | |
| mask_ratio_max: float = 1.0 | |
| dtype: torch.dtype = torch.float32 # type: ignore | |
| tie_word_embeddings: bool = False | |
| class _RotaryEmbedding(nn.Module): | |
| def __init__(self, dim, max_position_embeddings=5120, base=10000.0, device=None): | |
| super().__init__() | |
| self.dim = dim | |
| self.max_position_embeddings = max_position_embeddings | |
| self.base = base | |
| inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, device=device).float() / dim)) | |
| self.register_buffer("inv_freq", inv_freq, persistent=False) | |
| def forward(self, x, position_ids): | |
| inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) | |
| position_ids_expanded = position_ids[:, None, :].float() | |
| freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2) | |
| emb = torch.cat((freqs, freqs), dim=-1) | |
| cos = emb.cos() | |
| sin = emb.sin() | |
| return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) | |
| def _rotate_half(x): | |
| x1, x2 = x.chunk(2, dim=-1) | |
| return torch.cat((-x2, x1), dim=-1) | |
| def _apply_rotary_pos_emb(q, k, cos, sin): | |
| cos = cos.unsqueeze(1) | |
| sin = sin.unsqueeze(1) | |
| q_embed = (q * cos) + (_rotate_half(q) * sin) | |
| k_embed = (k * cos) + (_rotate_half(k) * sin) | |
| return q_embed, k_embed | |
| class _TimestepEmbedding(nn.Module): | |
| def __init__(self, hidden_size): | |
| super().__init__() | |
| self.hidden_size = hidden_size | |
| self.mlp = nn.Sequential( | |
| nn.Linear(hidden_size, hidden_size * 4), | |
| nn.SiLU(), | |
| nn.Linear(hidden_size * 4, hidden_size), | |
| ) | |
| def forward(self, t): | |
| half_dim = self.hidden_size // 2 | |
| emb = math.log(10000.0) / (half_dim - 1) | |
| emb = torch.exp(torch.arange(half_dim, device=t.device, dtype=torch.float32) * -emb) | |
| emb = t[:, None].float() * emb[None, :] | |
| emb = torch.cat([emb.sin(), emb.cos()], dim=-1) | |
| return self.mlp(emb).to(t.dtype) | |
| class _TimestepResidual(nn.Module): | |
| def __init__(self, hidden_size): | |
| super().__init__() | |
| self.proj = nn.Linear(hidden_size, hidden_size) | |
| nn.init.zeros_(self.proj.weight) | |
| nn.init.zeros_(self.proj.bias) | |
| def forward(self, x, emb): | |
| return x + self.proj(emb)[:, None, :] | |
| class _RMSNorm(nn.Module): | |
| def __init__(self, hidden_size, eps=1e-6): | |
| super().__init__() | |
| self.weight = nn.Parameter(torch.ones(hidden_size)) | |
| self.eps = eps | |
| def forward(self, x): | |
| var = x.pow(2).mean(-1, keepdim=True) | |
| x = x * torch.rsqrt(var + self.eps) | |
| return self.weight * x | |
| class _SelfAttention(nn.Module): | |
| def __init__(self, config: MetaDiffusionConfig): | |
| super().__init__() | |
| self.config = config | |
| self.hidden_size = config.hidden_size | |
| self.num_heads = config.num_attention_heads | |
| self.num_kv_heads = config.num_key_value_heads | |
| self.head_dim = config.head_dim | |
| self.num_kv_groups = self.num_heads // self.num_kv_heads | |
| self.q_proj = nn.Linear(config.hidden_size, self.num_heads * config.head_dim, bias=False) | |
| self.k_proj = nn.Linear(config.hidden_size, self.num_kv_heads * config.head_dim, bias=False) | |
| self.v_proj = nn.Linear(config.hidden_size, self.num_kv_heads * config.head_dim, bias=False) | |
| self.o_proj = nn.Linear(self.num_heads * config.head_dim, config.hidden_size, bias=False) | |
| self.rotary_emb = _RotaryEmbedding(config.head_dim, max_position_embeddings=config.max_position_embeddings, base=config.rope_theta) | |
| def forward(self, x, attention_mask=None, position_ids=None): | |
| batch, seq, _ = x.shape | |
| q = self.q_proj(x).view(batch, seq, self.num_heads, self.head_dim).transpose(1, 2) | |
| k = self.k_proj(x).view(batch, seq, self.num_kv_heads, self.head_dim).transpose(1, 2) | |
| v = self.v_proj(x).view(batch, seq, self.num_kv_heads, self.head_dim).transpose(1, 2) | |
| cos, sin = self.rotary_emb(x, position_ids) | |
| q, k = _apply_rotary_pos_emb(q, k, cos, sin) | |
| if self.num_kv_groups > 1: | |
| k = k.repeat_interleave(self.num_kv_groups, dim=1) | |
| v = v.repeat_interleave(self.num_kv_groups, dim=1) | |
| out = F.scaled_dot_product_attention(q, k, v, attn_mask=attention_mask) | |
| out = out.transpose(1, 2).contiguous().view(batch, seq, -1) | |
| return self.o_proj(out) | |
| class _MLP(nn.Module): | |
| def __init__(self, config: MetaDiffusionConfig): | |
| super().__init__() | |
| self.gate_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False) | |
| self.up_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False) | |
| self.down_proj = nn.Linear(config.intermediate_size, config.hidden_size, bias=False) | |
| def forward(self, x): | |
| return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) | |
| class _TransformerBlock(nn.Module): | |
| def __init__(self, config: MetaDiffusionConfig): | |
| super().__init__() | |
| self.input_layernorm = _RMSNorm(config.hidden_size, eps=config.rms_norm_eps) | |
| self.self_attn = _SelfAttention(config) | |
| self.post_attention_layernorm = _RMSNorm(config.hidden_size, eps=config.rms_norm_eps) | |
| self.mlp = _MLP(config) | |
| self.timestep_residual = _TimestepResidual(config.hidden_size) | |
| def forward(self, x, timestep_emb, attention_mask=None, position_ids=None): | |
| residual = x | |
| x = self.input_layernorm(x) | |
| x = self.self_attn(x, attention_mask, position_ids) | |
| x = residual + x | |
| x = self.timestep_residual(x, timestep_emb) | |
| residual = x | |
| x = self.post_attention_layernorm(x) | |
| x = self.mlp(x) | |
| x = residual + x | |
| x = self.timestep_residual(x, timestep_emb) | |
| return x | |
| class MetaDiffusionLM(nn.Module): | |
| def __init__(self, config: MetaDiffusionConfig): | |
| super().__init__() | |
| self.config = config | |
| self.embed_tokens = nn.Embedding(config.mask_vocab_size, config.hidden_size, padding_idx=config.pad_token_id) | |
| self.timestep_emb = _TimestepEmbedding(config.timestep_emb_hidden) | |
| self.layers = nn.ModuleList([_TransformerBlock(config) for _ in range(config.num_hidden_layers)]) | |
| self.norm = _RMSNorm(config.hidden_size, eps=config.rms_norm_eps) | |
| if config.tie_word_embeddings: | |
| self.lm_head = None # type: ignore | |
| else: | |
| self.lm_head = nn.Linear(config.hidden_size, config.mask_vocab_size, bias=False) | |
| if self.lm_head is not None: | |
| nn.init.normal_(self.lm_head.weight, std=0.02) | |
| def forward(self, input_ids, timesteps, attention_mask=None): | |
| batch, seq = input_ids.shape | |
| position_ids = torch.arange(seq, device=input_ids.device).unsqueeze(0).expand(batch, -1) | |
| x = self.embed_tokens(input_ids) | |
| t_emb = self.timestep_emb(timesteps) | |
| attn_mask = None | |
| if attention_mask is not None: | |
| attn_mask = ((1.0 - attention_mask[:, None, None, :].float()) * -1e9).to(x.dtype) | |
| for layer in self.layers: | |
| x = layer(x, t_emb, attn_mask, position_ids) | |
| x = self.norm(x) | |
| if self.lm_head is not None: | |
| logits = self.lm_head(x) | |
| else: | |
| logits = F.linear(x, self.embed_tokens.weight) | |
| return logits | |
| DIFF_MASK_ID = 32000 | |
| DIFF_CHAT_TOKENS = ["<|im_start|>", "<|im_end|>"] + [f"<|r{i}|>" for i in range(1, 8)] | |
| DIFF_IM_START, DIFF_IM_END = "<|im_start|>", "<|im_end|>" | |
| def _ensure_diff_chat_tokens(tokenizer): | |
| """Add ChatML + rainbow tokens if missing (base tokenizer case). Mirrors chat.py.""" | |
| if tokenizer.convert_tokens_to_ids(DIFF_IM_START) == tokenizer.unk_token_id: | |
| if len(tokenizer) == 32000: | |
| tokenizer.add_special_tokens({"additional_special_tokens": ["<|reserved|>"]}) | |
| tokenizer.add_special_tokens({"additional_special_tokens": DIFF_CHAT_TOKENS}) | |
| assert tokenizer.convert_tokens_to_ids(DIFF_IM_END) == 32002, "chat token ids wrong (collide with mask id 32000)" | |
| return tokenizer | |
| def _format_diff_messages(messages): | |
| parts = [] | |
| for m in messages: | |
| parts.append(f"{DIFF_IM_START}{m['role']}\n{m['content']}{DIFF_IM_END}") | |
| return "\n".join(parts) | |
| def _diff_cumulative_unmask_frac(i, N): | |
| return 0.5 * (1 - math.cos(math.pi * i / N)) | |
| def _diff_cut_response(tokens, tokenizer): | |
| """Cut at <|im_end|> or </s>; drop rainbow/pad. Mirrors chat.py.""" | |
| im_end_id = tokenizer.convert_tokens_to_ids(DIFF_IM_END) | |
| eos_id = tokenizer.eos_token_id | |
| rainbow_ids = {tokenizer.convert_tokens_to_ids(f"<|r{i}|>") for i in range(1, 8)} | |
| out = [] | |
| for t in tokens: | |
| if t == im_end_id or t == eos_id: | |
| break | |
| if t in rainbow_ids or t == tokenizer.pad_token_id: | |
| continue | |
| out.append(t) | |
| return out | |
| def _diff_generate_response(model, tokenizer, prompt_ids, gen_len, num_steps, temperature, repetition_penalty, device, stop_on_end=True): | |
| model.eval() | |
| total_len = prompt_ids.shape[1] + gen_len | |
| x = torch.full((1, total_len), DIFF_MASK_ID, device=device, dtype=torch.long) | |
| x[0, : prompt_ids.shape[1]] = prompt_ids | |
| mask_id = DIFF_MASK_ID | |
| im_end_id = tokenizer.convert_tokens_to_ids(DIFF_IM_END) | |
| eos_id = tokenizer.eos_token_id | |
| prompt_len = prompt_ids.shape[1] | |
| for i in range(num_steps): | |
| frac_now = _diff_cumulative_unmask_frac(i, num_steps) | |
| frac_next = _diff_cumulative_unmask_frac(i + 1, num_steps) | |
| n_masked = (x == mask_id).sum().item() | |
| n_total = int((frac_next - frac_now) * gen_len + 0.5) | |
| if i == num_steps - 1: | |
| n_unmask = n_masked | |
| else: | |
| n_unmask = max(n_total, 1) if n_masked > 0 else 0 | |
| t = 1.0 - frac_now | |
| logits = model(x, torch.full((1,), t, device=device)) | |
| logits[:, :, mask_id] = -1e9 | |
| if repetition_penalty != 1.0: | |
| for tok in x[0].unique(): | |
| ti = int(tok.item()) | |
| if 0 <= ti < logits.shape[-1]: | |
| logits[0, :, ti] = torch.where( | |
| logits[0, :, ti] < 0, | |
| logits[0, :, ti] * repetition_penalty, | |
| logits[0, :, ti] / repetition_penalty, | |
| ) | |
| mask_positions = x == mask_id | |
| if not mask_positions.any(): | |
| break | |
| mask_logits = logits[mask_positions] | |
| probs = F.softmax(mask_logits / max(0.1, temperature), dim=-1) | |
| sampled = torch.multinomial(probs, 1).squeeze(-1) | |
| mask_flat = mask_positions.nonzero(as_tuple=False) | |
| if n_unmask < int(mask_positions.sum().item()): | |
| fill_positions = mask_flat[:n_unmask] | |
| for idx, tok in zip(fill_positions, sampled[:n_unmask]): | |
| x[idx[0], idx[1]] = tok | |
| else: | |
| x[mask_positions] = sampled | |
| if stop_on_end and ((x[0, prompt_len:] == im_end_id).any() or (x[0, prompt_len:] == eos_id).any()): | |
| break | |
| return x | |
| # --------------------------------------------------------------------------- | |
| # ELO persistence | |
| # --------------------------------------------------------------------------- | |
| def init_elo_state() -> Dict[str, dict]: | |
| return {mid: {"rating": float(INIT_RATING), "wins": 0, "losses": 0, "battles": 0, "ties": 0, "both_bad": 0} for mid in MODEL_IDS} | |
| def load_elo() -> Dict[str, dict]: | |
| if get_elo_file().exists(): | |
| try: | |
| with open(get_elo_file(), "r") as f: | |
| data = json.load(f) | |
| for mid in MODEL_IDS: | |
| if mid not in data: | |
| data[mid] = {"rating": float(INIT_RATING), "wins": 0, "losses": 0, "battles": 0, "ties": 0, "both_bad": 0} | |
| else: | |
| data[mid].setdefault("rating", float(INIT_RATING)) | |
| data[mid].setdefault("wins", 0) | |
| data[mid].setdefault("losses", 0) | |
| data[mid].setdefault("battles", 0) | |
| data[mid].setdefault("ties", 0) | |
| data[mid].setdefault("both_bad", 0) | |
| return data | |
| except Exception as e: | |
| logger.warning(f"Failed to load ELO file: {e}, resetting") | |
| return init_elo_state() | |
| def save_elo(state: Dict[str, dict]): | |
| try: | |
| get_data_dir().mkdir(parents=True, exist_ok=True) | |
| with open(get_elo_file(), "w") as f: | |
| json.dump(state, f, indent=2) | |
| except Exception as e: | |
| logger.error(f"Failed to save ELO: {e}") | |
| def expected_score(ra: float, rb: float) -> float: | |
| return 1.0 / (1.0 + BASE ** ((rb - ra) / SCALE)) | |
| def update_elo(state: Dict[str, dict], model_a: str, model_b: str, winner: Optional[str]) -> Dict[str, dict]: | |
| if model_a not in state or model_b not in state: | |
| logger.warning(f"Unknown models in ELO update: {model_a}, {model_b}") | |
| return state | |
| ra = state[model_a]["rating"] | |
| rb = state[model_b]["rating"] | |
| ea = expected_score(ra, rb) | |
| eb = expected_score(rb, ra) | |
| if winner == model_a: | |
| sa = 1.0 | |
| elif winner == model_b: | |
| sa = 0.0 | |
| elif winner is None or winner == "tie" or winner == "both_bad": | |
| sa = 0.5 | |
| else: | |
| raise ValueError(f"Unexpected winner: {winner}") | |
| sb = 1.0 - sa | |
| state[model_a]["rating"] = ra + K_FACTOR * (sa - ea) | |
| state[model_b]["rating"] = rb + K_FACTOR * (sb - eb) | |
| state[model_a]["battles"] += 1 | |
| state[model_b]["battles"] += 1 | |
| if sa == 1.0: | |
| state[model_a]["wins"] += 1 | |
| state[model_b]["losses"] += 1 | |
| elif sa == 0.0: | |
| state[model_b]["wins"] += 1 | |
| state[model_a]["losses"] += 1 | |
| elif winner != "both_bad": | |
| state[model_a]["ties"] += 1 | |
| state[model_b]["ties"] += 1 | |
| if winner == "both_bad": | |
| state[model_a]["both_bad"] = state[model_a].get("both_bad", 0) + 1 | |
| state[model_b]["both_bad"] = state[model_b].get("both_bad", 0) + 1 | |
| save_elo(state) | |
| return state | |
| def leaderboard_dataframe(state: Optional[Dict[str, dict]] = None) -> pd.DataFrame: | |
| if state is None: | |
| state = load_elo() | |
| rows = [] | |
| for mid in MODEL_IDS: | |
| info = state.get(mid, {"rating": INIT_RATING, "wins": 0, "losses": 0, "battles": 0, "ties": 0}) | |
| rows.append({ | |
| "Model": MODEL_DISPLAY.get(mid, mid), | |
| "Model ID": mid, | |
| "ELO": round(float(info["rating"]), 1), | |
| "Battles": int(info["battles"]), | |
| "Wins": int(info["wins"]), | |
| "Losses": int(info["losses"]), | |
| "Ties": int(info.get("ties", 0)), | |
| "Both Bad": int(info.get("both_bad", 0)), | |
| }) | |
| df = pd.DataFrame(rows) | |
| df = df.sort_values(by="ELO", ascending=False).reset_index(drop=True) | |
| df.insert(0, "Rank", range(1, len(df) + 1)) | |
| return df | |
| # --------------------------------------------------------------------------- | |
| # Chat logging to data/chats.jsonl | |
| # --------------------------------------------------------------------------- | |
| def log_battle(prompt: str, model_a: str, model_b: str, response_a: str, response_b: str, chosen: str, winner_model: str): | |
| """ | |
| Append one battle record to data/chats.jsonl. | |
| Fields: prompt, response_a, response_b, model_a, model_b, chosen (A/B/tie/both_bad), winner_model, timestamp | |
| Spec: keeps user's message, two responses, each model's names, and what response user chose. | |
| """ | |
| try: | |
| get_data_dir().mkdir(parents=True, exist_ok=True) | |
| record = { | |
| "timestamp": __import__("datetime").datetime.now(__import__("datetime").timezone.utc).isoformat(), | |
| "prompt": prompt, | |
| "model_a": model_a, | |
| "model_b": model_b, | |
| "response_a": response_a, | |
| "response_b": response_b, | |
| "chosen": chosen, # "A" / "B" / "tie" / "both_bad" | |
| "winner_model": winner_model, | |
| "chosen_response": response_a if chosen == "A" else response_b if chosen == "B" else "", | |
| } | |
| with open(get_chat_file(), "a", encoding="utf-8") as f: | |
| f.write(json.dumps(record, ensure_ascii=False) + "\n") | |
| except Exception as e: | |
| logger.error(f"Failed to log battle: {e}") | |
| # --------------------------------------------------------------------------- | |
| # Model loading (CPU) | |
| # --------------------------------------------------------------------------- | |
| models: Dict[str, object] = {} | |
| tokenizers: Dict[str, object] = {} | |
| model_load_errors: Dict[str, str] = {} | |
| # Diffusion manual instance (if loaded) | |
| diffusion_model: Optional[MetaDiffusionLM] = None | |
| diffusion_tokenizer = None | |
| HF_DIFFUSION_REPO = "CodeSoft/MetaDiffusion-150M-ChatBase" | |
| def load_diffusion_manual(): | |
| """Load MetaDiffusion from HuggingFace (only) using inline architecture.""" | |
| global diffusion_model, diffusion_tokenizer | |
| if diffusion_model is not None: | |
| # Re-register in global dicts if cleared (e.g., after tests) | |
| if "CodeSoft/MetaDiffusion-150M-ChatBase" not in models: | |
| models["CodeSoft/MetaDiffusion-150M-ChatBase"] = diffusion_model # type: ignore | |
| if diffusion_tokenizer is not None and "CodeSoft/MetaDiffusion-150M-ChatBase" not in tokenizers: | |
| tokenizers["CodeSoft/MetaDiffusion-150M-ChatBase"] = diffusion_tokenizer # type: ignore | |
| return diffusion_model, diffusion_tokenizer | |
| try: | |
| from huggingface_hub import snapshot_download | |
| repo_id = HF_DIFFUSION_REPO | |
| local_dir = Path(snapshot_download(repo_id)) | |
| cfg_path = local_dir / "config.json" | |
| tok_path = local_dir | |
| model_path = local_dir / "model.safetensors" | |
| if not cfg_path.exists() or not model_path.exists(): | |
| logger.warning(f"Diffusion files not found in HF snapshot {local_dir}") | |
| return None, None | |
| with open(cfg_path, "r") as f: | |
| cfg_dict = json.load(f) | |
| valid = {k: v for k, v in cfg_dict.items() if k in MetaDiffusionConfig.__dataclass_fields__} | |
| cfg = MetaDiffusionConfig(**valid) | |
| cfg.tie_word_embeddings = False | |
| mdl = MetaDiffusionLM(cfg).to(DEVICE) | |
| try: | |
| from safetensors.torch import load_file | |
| except ImportError: | |
| import subprocess, sys | |
| subprocess.check_call([sys.executable, "-m", "pip", "install", "safetensors", "--quiet", "--break-system-packages"]) | |
| from safetensors.torch import load_file # type: ignore | |
| state = load_file(str(model_path), device="cpu") | |
| state = {k[len("model."):] if k.startswith("model.") else k: v for k, v in state.items()} | |
| missing, unexpected = mdl.load_state_dict(state, strict=False) | |
| if missing or unexpected: | |
| logger.info(f" Diffusion load: missing={missing[:3]} unexpected={unexpected[:3]}") | |
| mdl.to(DEVICE) | |
| mdl.eval() | |
| logger.info(f" Loaded {sum(p.numel() for p in mdl.parameters())/1e6:.1f}M params, vocab={cfg.mask_vocab_size}") | |
| tok = AutoTokenizer.from_pretrained(str(tok_path), trust_remote_code=True) | |
| tok = _ensure_diff_chat_tokens(tok) | |
| if tok.pad_token is None: | |
| tok.pad_token = tok.eos_token | |
| diffusion_model = mdl | |
| diffusion_tokenizer = tok | |
| logger.info(f"[+] Loaded MetaDiffusion manual from HF {repo_id} (vocab {len(tok)})") | |
| models["CodeSoft/MetaDiffusion-150M-ChatBase"] = mdl # type: ignore | |
| tokenizers["CodeSoft/MetaDiffusion-150M-ChatBase"] = tok # type: ignore | |
| return mdl, tok | |
| except Exception as e: | |
| logger.warning(f"Manual diffusion load failed: {e}\n{traceback.format_exc()}") | |
| return None, None | |
| LOCAL_PATHS: Dict[str, str] = {} | |
| def load_models(): | |
| global models, tokenizers, model_load_errors | |
| # If already populated (including diffusion manual), return | |
| # But we want to ensure all 5 attempted | |
| if models and len(models) >= 3: | |
| # Already loaded, but ensure diffusion tried | |
| if "CodeSoft/MetaDiffusion-150M-ChatBase" not in models: | |
| load_diffusion_manual() | |
| return models, tokenizers | |
| logger.info(f"Loading {len(MODEL_IDS)} models on {DEVICE} ...") | |
| # Try diffusion manual first (bypass HF Auto which fails on unknown type) | |
| if "CodeSoft/MetaDiffusion-150M-ChatBase" not in models: | |
| load_diffusion_manual() | |
| for mid in MODEL_IDS: | |
| if mid in models: | |
| continue # already loaded (diffusion) | |
| load_id = LOCAL_PATHS.get(mid, mid) if os.path.exists(LOCAL_PATHS.get(mid, "")) else mid | |
| candidates = [load_id] | |
| if mid in FALLBACK_IDS: | |
| candidates.append(FALLBACK_IDS[mid]) | |
| success = False | |
| last_err = None | |
| for cand in candidates: | |
| try: | |
| logger.info(f"[*] Loading {mid} (candidate {cand})...") | |
| tok = AutoTokenizer.from_pretrained(cand, trust_remote_code=True) | |
| if tok.pad_token is None: | |
| tok.pad_token = tok.eos_token | |
| mdl = AutoModelForCausalLM.from_pretrained( | |
| cand, | |
| trust_remote_code=True, | |
| torch_dtype=torch.float32, | |
| low_cpu_mem_usage=True, | |
| ) | |
| mdl.to(DEVICE) | |
| mdl.eval() | |
| tokenizers[mid] = tok | |
| models[mid] = mdl | |
| logger.info(f"[+] Loaded {mid} from {cand} (tok vocab {len(tok)})") | |
| success = True | |
| break | |
| except Exception as e: | |
| last_err = f"{e}\n{traceback.format_exc()}" | |
| logger.warning(f"Failed to load {mid} from {cand}: {e}") | |
| continue | |
| if not success: | |
| err_msg = f"Failed candidates {candidates}: {last_err}" | |
| model_load_errors[mid] = err_msg | |
| logger.warning(f"[!] {mid} failed to load — generation will error. Error: {err_msg[:600]}") | |
| logger.info(f"Model loading complete. Loaded: {list(models.keys())} | Failed: {list(model_load_errors.keys())}") | |
| return models, tokenizers | |
| def ensure_models_loaded(): | |
| # Load if not already attempted | |
| if not models and not model_load_errors: | |
| load_models() | |
| elif "CodeSoft/MetaDiffusion-150M-ChatBase" not in models and not model_load_errors.get("CodeSoft/MetaDiffusion-150M-ChatBase"): | |
| # Try diffusion again if not yet loaded | |
| load_diffusion_manual() | |
| # --------------------------------------------------------------------------- | |
| # Prompt formatting & generation | |
| # --------------------------------------------------------------------------- | |
| def build_inputs(tokenizer, model_id: str, prompt: str): | |
| ctx = MODEL_CONTEXT.get(model_id, 2048) | |
| gen_budget = GEN_DEFAULTS.get(model_id, {}).get("max_new_tokens", 128) | |
| max_prompt_tokens = max(32, ctx - gen_budget - 16) | |
| try: | |
| if hasattr(tokenizer, "chat_template") and tokenizer.chat_template is not None: | |
| messages = [{"role": "user", "content": prompt}] | |
| inputs = tokenizer.apply_chat_template( | |
| messages, add_generation_prompt=True, return_tensors="pt", truncation=True, max_length=max_prompt_tokens | |
| ) | |
| if isinstance(inputs, torch.Tensor): | |
| inputs = {"input_ids": inputs} | |
| for k in list(inputs.keys()): | |
| if isinstance(inputs[k], torch.Tensor): | |
| inputs[k] = inputs[k].to(DEVICE) | |
| return inputs | |
| elif hasattr(tokenizer, "apply_chat_template"): | |
| try: | |
| messages = [{"role": "user", "content": prompt}] | |
| inputs = tokenizer.apply_chat_template( | |
| messages, add_generation_prompt=True, return_tensors="pt", truncation=True, max_length=max_prompt_tokens | |
| ) | |
| if isinstance(inputs, torch.Tensor): | |
| inputs = {"input_ids": inputs} | |
| for k in list(inputs.keys()): | |
| if isinstance(inputs[k], torch.Tensor): | |
| inputs[k] = inputs[k].to(DEVICE) | |
| return inputs | |
| except Exception: | |
| pass | |
| except Exception as e: | |
| logger.debug(f"Chat template failed for {model_id}: {e}") | |
| inputs = tokenizer(prompt, return_tensors="pt", truncation=True, max_length=max_prompt_tokens) | |
| for k in list(inputs.keys()): | |
| if isinstance(inputs[k], torch.Tensor): | |
| inputs[k] = inputs[k].to(DEVICE) | |
| return inputs | |
| def is_diffusion_model(model_id: str) -> bool: | |
| return "metadiffusion" in model_id.lower() | |
| def generate_for_model(model_id: str, prompt: str) -> str: | |
| ensure_models_loaded() | |
| if model_id not in models or model_id not in tokenizers: | |
| short = MODEL_DISPLAY.get(model_id, model_id) | |
| err = model_load_errors.get(model_id, "model not loaded") | |
| err_short = str(err).splitlines()[0][:800] if err else "model not loaded" | |
| return f"[Error: {model_id} not loaded: {err_short}]" | |
| tokenizer = tokenizers[model_id] | |
| model = models[model_id] | |
| cfg = GEN_DEFAULTS.get(model_id, {}) | |
| max_new = cfg.get("max_new_tokens", 128) | |
| try: | |
| if is_diffusion_model(model_id): | |
| return generate_diffusion(model, tokenizer, prompt, cfg) # type: ignore | |
| inputs = build_inputs(tokenizer, model_id, prompt) | |
| input_len = inputs["input_ids"].shape[1] | |
| gen_kwargs = { | |
| "max_new_tokens": max_new, | |
| "do_sample": cfg.get("do_sample", True), | |
| "temperature": cfg.get("temperature", 0.7), | |
| "top_p": cfg.get("top_p", 0.9), | |
| "repetition_penalty": cfg.get("repetition_penalty", 1.1), | |
| "pad_token_id": tokenizer.pad_token_id if tokenizer.pad_token_id is not None else tokenizer.eos_token_id, | |
| "eos_token_id": tokenizer.eos_token_id, | |
| "use_cache": False, | |
| } | |
| if "top_k" in cfg: | |
| gen_kwargs["top_k"] = cfg["top_k"] | |
| if "no_repeat_ngram_size" in cfg: | |
| gen_kwargs["no_repeat_ngram_size"] = cfg["no_repeat_ngram_size"] | |
| ctx = MODEL_CONTEXT.get(model_id, 2048) | |
| if input_len + max_new > ctx: | |
| gen_kwargs["max_new_tokens"] = max(16, ctx - input_len - 4) | |
| with torch.inference_mode(): | |
| outputs = model.generate(**inputs, **gen_kwargs) # type: ignore | |
| new_tokens = outputs[0, input_len:] | |
| text = tokenizer.decode(new_tokens, skip_special_tokens=True).strip() | |
| if not text: | |
| text = tokenizer.decode(outputs[0], skip_special_tokens=True).strip() | |
| prompt_text = tokenizer.decode(inputs["input_ids"][0], skip_special_tokens=True).strip() | |
| if text.startswith(prompt_text): | |
| text = text[len(prompt_text):].strip() | |
| return text if text else "[Empty response]" | |
| except Exception as e: | |
| logger.error(f"Generation failed for {model_id}: {e}\n{traceback.format_exc()}") | |
| return f"[Error generating from {MODEL_DISPLAY.get(model_id, model_id)}: {str(e)[:200]}]" | |
| def generate_diffusion(model, tokenizer, prompt: str, cfg: dict) -> str: | |
| try: | |
| tokenizer = _ensure_diff_chat_tokens(tokenizer) | |
| messages = [{"role": "user", "content": prompt}] | |
| prompt_str = _format_diff_messages(messages) + f"\n{DIFF_IM_START}assistant\n" | |
| prompt_ids = torch.tensor([tokenizer.encode(prompt_str, add_special_tokens=False)], device=DEVICE) | |
| gen_len = int(cfg.get("max_new_tokens", 96)) | |
| num_steps = int(cfg.get("num_steps", 128)) | |
| temperature = float(cfg.get("temperature", 0.7)) | |
| repetition_penalty = float(cfg.get("repetition_penalty", 1.5)) | |
| max_ctx = MODEL_CONTEXT.get("CodeSoft/MetaDiffusion-150M-ChatBase", 5120) | |
| if prompt_ids.shape[1] + gen_len > max_ctx: | |
| gen_len = max(16, max_ctx - prompt_ids.shape[1] - 4) | |
| if gen_len > 256: | |
| gen_len = 256 | |
| for attempt in range(3): | |
| cur_temp = temperature * (1 + 0.15 * attempt) | |
| x = _diff_generate_response( | |
| model, tokenizer, prompt_ids, gen_len, num_steps, cur_temp, repetition_penalty, DEVICE, stop_on_end=True | |
| ) | |
| response_tokens = x[0, prompt_ids.shape[1]:].tolist() | |
| response_tokens = _diff_cut_response(response_tokens, tokenizer) | |
| text = tokenizer.decode(response_tokens, skip_special_tokens=True).strip() | |
| if text: | |
| return text | |
| return "(empty response)" | |
| except Exception as e: | |
| logger.warning(f"Diffusion chat failed: {e}\n{traceback.format_exc()}") | |
| return f"[Diffusion error] {str(e)[:200]}" | |
| # --------------------------------------------------------------------------- | |
| # Gradio UI | |
| # --------------------------------------------------------------------------- | |
| CSS = """ | |
| .gradio-container {max-width: 1450px !important; width: 95% !important;} | |
| .vote-btn {font-weight: 700 !important;} | |
| /* Leaderboard: prevent ELO wrapping, give it fixed width */ | |
| #leaderboard { overflow-x: auto; } | |
| #leaderboard table { table-layout: auto; width: 100%; } | |
| #leaderboard th:nth-child(4), #leaderboard td:nth-child(4) { | |
| min-width: 95px; | |
| width: 95px; | |
| white-space: nowrap; | |
| text-align: center; | |
| font-variant-numeric: tabular-nums; | |
| } | |
| #leaderboard th:nth-child(1), #leaderboard td:nth-child(1) { min-width: 55px; width: 55px; text-align: center; } | |
| #leaderboard td { white-space: nowrap; overflow: hidden; text-overflow: ellipsis; } | |
| """ | |
| def pick_random_pair(exclude_pair: Optional[Tuple[str, str]] = None) -> Tuple[str, str]: | |
| state = load_elo() | |
| models_list = MODEL_IDS[:] | |
| weights = [] | |
| C = 5 | |
| K = 100 | |
| for m in models_list: | |
| games = state.get(m, {}).get("battles", 0) | |
| w = K / (games + C) | |
| weights.append(w) | |
| a = random.choices(models_list, weights=weights, k=1)[0] | |
| remaining = [m for m in models_list if m != a] | |
| remaining_weights = [w for m, w in zip(models_list, weights) if m != a] | |
| b = random.choices(remaining, weights=remaining_weights, k=1)[0] | |
| if exclude_pair and set((a, b)) == set(exclude_pair): | |
| a, b = random.sample(MODEL_IDS, 2) | |
| return a, b | |
| def create_demo() -> gr.Blocks: | |
| state_init = load_elo() | |
| df_init = leaderboard_dataframe(state_init) | |
| with gr.Blocks(title="SLM Arena") as demo: | |
| gr.Markdown( | |
| """ | |
| # ⚔️ SLM Arena | |
| """ | |
| ) | |
| last_pair = gr.State(None) | |
| with gr.Tabs(): | |
| with gr.Tab("Arena", id=0): | |
| prompt = gr.Textbox( | |
| label="Your prompt", | |
| placeholder="Ask anything... e.g. 'Explain quantum computing in simple terms' or 'Write a haiku about rain'", | |
| lines=3, | |
| ) | |
| with gr.Row(): | |
| submit_btn = gr.Button("⚔️ Battle", variant="primary", scale=1) | |
| clear_btn = gr.Button("Clear", variant="secondary", scale=1) | |
| with gr.Row(): | |
| with gr.Column(): | |
| response_a = gr.Textbox( | |
| label="Model A", lines=10, max_lines=14, interactive=False, | |
| placeholder="Response A will appear here..." | |
| ) | |
| reveal_a = gr.Markdown(visible=False) | |
| with gr.Column(): | |
| response_b = gr.Textbox( | |
| label="Model B", lines=10, max_lines=14, interactive=False, | |
| placeholder="Response B will appear here..." | |
| ) | |
| reveal_b = gr.Markdown(visible=False) | |
| with gr.Row(): | |
| vote_a = gr.Button("👈 Vote for A", variant="secondary", interactive=False, elem_classes=["vote-btn"]) | |
| vote_tie = gr.Button("🤝 Tie", variant="secondary", interactive=False, elem_classes=["vote-btn"]) | |
| vote_both_bad = gr.Button("👎 Both Bad", variant="secondary", interactive=False, elem_classes=["vote-btn"]) | |
| vote_b = gr.Button("Vote for B 👉", variant="secondary", interactive=False, elem_classes=["vote-btn"]) | |
| status = gr.Markdown(visible=False) | |
| new_round_btn = gr.Button("🔄 New Round", visible=False, variant="secondary") | |
| model_a_state = gr.State("") | |
| model_b_state = gr.State("") | |
| voted_state = gr.State(False) | |
| prompt_state = gr.State("") | |
| leaderboard_tab = gr.Tab("Leaderboard", id=1) | |
| with leaderboard_tab: | |
| gr.Markdown("### 🏆 ELO Leaderboard") | |
| leaderboard = gr.Dataframe( | |
| value=df_init, | |
| headers=["Rank", "Model", "Model ID", "ELO", "Battles", "Wins", "Losses", "Ties", "Both Bad"], | |
| datatype=["number", "str", "str", "number", "number", "number", "number", "number", "number"], | |
| interactive=False, | |
| wrap=False, | |
| column_widths=["5%", "15%", "25%", "12%", "7%", "7%", "7%", "7%", "7%"], | |
| elem_id="leaderboard", | |
| ) | |
| with gr.Row(): | |
| refresh_btn = gr.Button("🔄 Refresh", variant="secondary") | |
| # ------------------------------------------------------------------- | |
| # Event handlers | |
| # ------------------------------------------------------------------- | |
| def on_submit(user_prompt: str, last_pair_val): | |
| user_prompt = (user_prompt or "").strip() | |
| if not user_prompt: | |
| return ( | |
| gr.update(value="", placeholder="Please enter a prompt first!"), | |
| gr.update(value=""), | |
| gr.update(visible=False), | |
| gr.update(visible=False), | |
| gr.update(visible=False, value=""), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(visible=False), | |
| "", "", False, user_prompt, last_pair_val, | |
| leaderboard_dataframe(load_elo()) | |
| ) | |
| a, b = pick_random_pair(exclude_pair=last_pair_val) | |
| if random.random() < 0.5: | |
| a, b = b, a | |
| ensure_models_loaded() | |
| resp_a = generate_for_model(a, user_prompt) | |
| resp_b = generate_for_model(b, user_prompt) | |
| if not resp_a.strip(): | |
| resp_a = "[No output... model returned empty]" | |
| if not resp_b.strip(): | |
| resp_b = "[No output... model returned empty]" | |
| return ( | |
| gr.update(value=resp_a), | |
| gr.update(value=resp_b), | |
| gr.update(visible=False), | |
| gr.update(visible=False), | |
| gr.update(visible=False, value=""), | |
| gr.update(interactive=True), | |
| gr.update(interactive=True), | |
| gr.update(interactive=True), | |
| gr.update(interactive=True), | |
| gr.update(visible=False), | |
| a, b, False, user_prompt, (a, b), | |
| leaderboard_dataframe(load_elo()) | |
| ) | |
| def on_vote(choice: str, model_a: str, model_b: str, resp_a: str, resp_b: str, user_prompt: str, voted: bool): | |
| if voted or not model_a or not model_b: | |
| return ( | |
| gr.update(visible=False), | |
| gr.update(visible=False), | |
| gr.update(visible=False, value=""), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(visible=False), | |
| voted, | |
| leaderboard_dataframe(load_elo()) | |
| ) | |
| if choice == "A": | |
| winner = model_a | |
| win_label = "A" | |
| chosen = "A" | |
| elif choice == "B": | |
| winner = model_b | |
| win_label = "B" | |
| chosen = "B" | |
| elif choice == "Tie": | |
| winner = None | |
| win_label = "Tie" | |
| chosen = "tie" | |
| elif choice == "Both Bad": | |
| winner = "both_bad" | |
| win_label = "Both Bad" | |
| chosen = "both_bad" | |
| else: | |
| winner = model_b | |
| win_label = "B" | |
| chosen = "B" | |
| state = load_elo() | |
| ra_before = state[model_a]["rating"] | |
| rb_before = state[model_b]["rating"] | |
| update_elo(state, model_a, model_b, winner) | |
| ra_after = state[model_a]["rating"] | |
| rb_after = state[model_b]["rating"] | |
| delta_a = ra_after - ra_before | |
| delta_b = rb_after - rb_before | |
| reveal_a_text = f"**Model A:** `{model_a}` ({MODEL_DISPLAY.get(model_a, model_a)}) — ELO {ra_after:.1f} ({delta_a:+.1f})" | |
| reveal_b_text = f"**Model B:** `{model_b}` ({MODEL_DISPLAY.get(model_b, model_b)}) — ELO {rb_after:.1f} ({delta_b:+.1f})" | |
| if choice == "Tie": | |
| status_text = ( | |
| f"You voted **Tie**: no winner\n\n" | |
| f"**ELO update:** {MODEL_DISPLAY.get(model_a, model_a)} {ra_before:.1f} → {ra_after:.1f} ({delta_a:+.1f}) | " | |
| f"{MODEL_DISPLAY.get(model_b, model_b)} {rb_before:.1f} → {rb_after:.1f} ({delta_b:+.1f})" | |
| ) | |
| elif choice == "Both Bad": | |
| status_text = ( | |
| f"You voted **Both Bad**: no winner\n\n" | |
| f"**ELO update:** {MODEL_DISPLAY.get(model_a, model_a)} {ra_before:.1f} → {ra_after:.1f} ({delta_a:+.1f}) | " | |
| f"{MODEL_DISPLAY.get(model_b, model_b)} {rb_before:.1f} → {rb_after:.1f} ({delta_b:+.1f})" | |
| ) | |
| else: | |
| status_text = ( | |
| f"You voted **{win_label}**: the winner is `{winner}`\n\n" | |
| f"**ELO update:** {MODEL_DISPLAY.get(model_a, model_a)} {ra_before:.1f} → {ra_after:.1f} ({delta_a:+.1f}) | " | |
| f"{MODEL_DISPLAY.get(model_b, model_b)} {rb_before:.1f} → {rb_after:.1f} ({delta_b:+.1f})" | |
| ) | |
| # Log chat to data/chats.jsonl | |
| log_battle(user_prompt, model_a, model_b, resp_a, resp_b, chosen, winner) | |
| df = leaderboard_dataframe(state) | |
| return ( | |
| gr.update(value=reveal_a_text, visible=True), | |
| gr.update(value=reveal_b_text, visible=True), | |
| gr.update(value=status_text, visible=True), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(visible=True), | |
| True, | |
| df | |
| ) | |
| def on_new_round(): | |
| return ( | |
| gr.update(value=""), | |
| gr.update(value=""), | |
| gr.update(value="", visible=False), | |
| gr.update(value="", visible=False), | |
| gr.update(value="", visible=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(visible=False), | |
| "", "", False, "" | |
| ) | |
| def on_clear(): | |
| return ( | |
| gr.update(value=""), | |
| gr.update(value=""), | |
| gr.update(value=""), | |
| gr.update(value="", visible=False), | |
| gr.update(value="", visible=False), | |
| gr.update(value="", visible=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(interactive=False), | |
| gr.update(visible=False), | |
| "", "", False, "" | |
| ) | |
| def on_refresh(): | |
| return leaderboard_dataframe(load_elo()) | |
| submit_btn.click( | |
| fn=on_submit, | |
| inputs=[prompt, last_pair], | |
| outputs=[response_a, response_b, reveal_a, reveal_b, status, vote_a, vote_tie, vote_both_bad, vote_b, new_round_btn, model_a_state, model_b_state, voted_state, prompt_state, last_pair, leaderboard], | |
| ) | |
| prompt.submit( | |
| fn=on_submit, | |
| inputs=[prompt, last_pair], | |
| outputs=[response_a, response_b, reveal_a, reveal_b, status, vote_a, vote_tie, vote_both_bad, vote_b, new_round_btn, model_a_state, model_b_state, voted_state, prompt_state, last_pair, leaderboard], | |
| ) | |
| vote_a.click( | |
| fn=lambda ma, mb, ra, rb, pr, vd: on_vote("A", ma, mb, ra, rb, pr, vd), | |
| inputs=[model_a_state, model_b_state, response_a, response_b, prompt_state, voted_state], | |
| outputs=[reveal_a, reveal_b, status, vote_a, vote_tie, vote_both_bad, vote_b, new_round_btn, voted_state, leaderboard], | |
| ) | |
| vote_tie.click( | |
| fn=lambda ma, mb, ra, rb, pr, vd: on_vote("Tie", ma, mb, ra, rb, pr, vd), | |
| inputs=[model_a_state, model_b_state, response_a, response_b, prompt_state, voted_state], | |
| outputs=[reveal_a, reveal_b, status, vote_a, vote_tie, vote_both_bad, vote_b, new_round_btn, voted_state, leaderboard], | |
| ) | |
| vote_both_bad.click( | |
| fn=lambda ma, mb, ra, rb, pr, vd: on_vote("Both Bad", ma, mb, ra, rb, pr, vd), | |
| inputs=[model_a_state, model_b_state, response_a, response_b, prompt_state, voted_state], | |
| outputs=[reveal_a, reveal_b, status, vote_a, vote_tie, vote_both_bad, vote_b, new_round_btn, voted_state, leaderboard], | |
| ) | |
| vote_b.click( | |
| fn=lambda ma, mb, ra, rb, pr, vd: on_vote("B", ma, mb, ra, rb, pr, vd), | |
| inputs=[model_a_state, model_b_state, response_a, response_b, prompt_state, voted_state], | |
| outputs=[reveal_a, reveal_b, status, vote_a, vote_tie, vote_both_bad, vote_b, new_round_btn, voted_state, leaderboard], | |
| ) | |
| new_round_btn.click( | |
| fn=on_new_round, | |
| inputs=[], | |
| outputs=[response_a, response_b, reveal_a, reveal_b, status, vote_a, vote_tie, vote_both_bad, vote_b, new_round_btn, model_a_state, model_b_state, voted_state, prompt_state], | |
| ) | |
| clear_btn.click( | |
| fn=on_clear, | |
| inputs=[], | |
| outputs=[prompt, response_a, response_b, reveal_a, reveal_b, status, vote_a, vote_tie, vote_both_bad, vote_b, new_round_btn, model_a_state, model_b_state, voted_state, prompt_state], | |
| ) | |
| refresh_btn.click(fn=on_refresh, inputs=[], outputs=[leaderboard]) | |
| # Refresh when Leaderboard tab is selected (fixes stale df_init) | |
| # Also refresh on page load but without global spinner (demo.load caused "loading..." until refresh when bucket slow) | |
| try: | |
| leaderboard_tab.select(fn=on_refresh, inputs=[], outputs=[leaderboard]) | |
| except Exception: | |
| pass | |
| # Page-load refresh without blocking UI (hidden progress) | |
| try: | |
| demo.load(fn=on_refresh, inputs=[], outputs=[leaderboard], show_progress="hidden") | |
| except Exception: | |
| # Fallback: no page-load auto-refresh, rely on tab select + initial df_init (now dynamic via get_data_dir) | |
| pass | |
| return demo | |
| # --------------------------------------------------------------------------- | |
| # Main | |
| # --------------------------------------------------------------------------- | |
| if __name__ == "__main__": | |
| print("=" * 60) | |
| print("SLM Arena starting, attempting to load 4 models on CPU...") | |
| print(f"Models: {MODEL_IDS}") | |
| print(f"Data dir: {get_data_dir().resolve()} (bucket /data if mounted)") | |
| print("=" * 60) | |
| try: | |
| load_models() | |
| except Exception as e: | |
| logger.error(f"Model loading encountered error: {e}") | |
| try: | |
| df = leaderboard_dataframe(load_elo()) | |
| print(df.to_string(index=False)) | |
| print(f"\nChat log: {get_chat_file().resolve()} (exists={get_chat_file().exists()})") | |
| if get_chat_file().exists(): | |
| with open(get_chat_file()) as f: | |
| lines = sum(1 for _ in f) | |
| print(f"Previous battles logged: {lines}") | |
| except Exception as e: | |
| logger.warning(f"Leaderboard preview failed: {e}") | |
| demo = create_demo() | |
| demo.queue(max_size=20) | |
| demo.launch(server_name="0.0.0.0", server_port=7860, show_error=True, theme=gr.themes.Base(), css=CSS) |