Files
model-eval-site/database.py
T

351 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""SQLite 存储:账号 / 评测提交 / 提交明细"""
import os
import json
import sqlite3
import threading
import config
_lock = threading.RLock()
SCHEMA = """
CREATE TABLE IF NOT EXISTS accounts(
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT UNIQUE NOT NULL,
remark TEXT DEFAULT '',
created_at TEXT DEFAULT (datetime('now','localtime'))
);
CREATE TABLE IF NOT EXISTS submissions(
id INTEGER PRIMARY KEY AUTOINCREMENT,
account_id INTEGER NOT NULL,
source_test_id INTEGER DEFAULT 0,
source_site TEXT DEFAULT 'llm-speed-tester',
created_at TEXT DEFAULT (datetime('now','localtime')),
provider TEXT DEFAULT '',
model TEXT DEFAULT '',
test_name TEXT DEFAULT '',
samples_ok INTEGER DEFAULT 0,
samples_total INTEGER DEFAULT 0,
avg_ttft_ms REAL,
avg_prefill_speed REAL,
avg_decode_speed REAL,
avg_stream_decode REAL,
avg_prompt_tokens REAL,
avg_output_tokens REAL,
avg_total_ms REAL,
min_ttft_ms REAL, max_ttft_ms REAL,
min_prefill_speed REAL, max_prefill_speed REAL,
min_decode_speed REAL, max_decode_speed REAL,
min_total_ms REAL, max_total_ms REAL,
context_lengths TEXT DEFAULT '[]',
concurrency_levels TEXT DEFAULT '[]',
by_length TEXT DEFAULT '{}',
by_concurrency TEXT DEFAULT '{}',
runs TEXT DEFAULT '[]',
raw TEXT DEFAULT '{}'
);
CREATE INDEX IF NOT EXISTS idx_sub_model ON submissions(model, provider);
CREATE INDEX IF NOT EXISTS idx_sub_account ON submissions(account_id);
CREATE INDEX IF NOT EXISTS idx_sub_created ON submissions(created_at);
"""
def _connect():
os.makedirs(config.DATA_DIR, exist_ok=True)
conn = sqlite3.connect(config.DB_PATH, check_same_thread=False, timeout=30)
conn.row_factory = sqlite3.Row
conn.execute("PRAGMA journal_mode=WAL")
return conn
def init_db():
with _lock:
conn = _connect()
try:
conn.executescript(SCHEMA)
conn.commit()
finally:
conn.close()
# ───────────────────────── 账号 ─────────────────────────
def get_or_create_account(name: str, remark: str = "") -> int:
name = (name or "").strip() or "默认账号"
with _lock:
conn = _connect()
try:
r = conn.execute("SELECT id FROM accounts WHERE name=?", (name,)).fetchone()
if r:
return r["id"]
cur = conn.execute("INSERT INTO accounts(name, remark) VALUES(?,?)",
(name, (remark or "").strip()))
conn.commit()
return cur.lastrowid
finally:
conn.close()
def add_account(name: str, remark: str = "") -> int:
name = (name or "").strip()
if not name:
raise ValueError("账号名称不能为空")
with _lock:
conn = _connect()
try:
r = conn.execute("SELECT id FROM accounts WHERE name=?", (name,)).fetchone()
if r:
raise ValueError("账号已存在")
cur = conn.execute("INSERT INTO accounts(name, remark) VALUES(?,?)", (name, remark))
conn.commit()
return cur.lastrowid
finally:
conn.close()
def list_accounts():
with _lock:
conn = _connect()
try:
rows = conn.execute(
"SELECT a.id,a.name,a.remark,a.created_at,"
"(SELECT COUNT(*) FROM submissions s WHERE s.account_id=a.id) AS cnt "
"FROM accounts a ORDER BY cnt DESC, a.id").fetchall()
return [dict(r) for r in rows]
finally:
conn.close()
def rename_account(aid: int, name: str, remark: str):
with _lock:
conn = _connect()
try:
conn.execute("UPDATE accounts SET name=?, remark=? WHERE id=?",
(name.strip(), (remark or "").strip(), aid))
conn.commit()
finally:
conn.close()
def delete_account(aid: int):
with _lock:
conn = _connect()
try:
conn.execute("DELETE FROM submissions WHERE account_id=?", (aid,))
conn.execute("DELETE FROM accounts WHERE id=?", (aid,))
conn.commit()
finally:
conn.close()
# ───────────────────────── 提交 ─────────────────────────
def add_submission(account_id: int, data: dict) -> int:
"""data 是 llm-speed-tester 发来的整包(含 summary/gen/runs/config"""
summary = data.get("summary") or {}
gen = data.get("gen") or {}
runs = data.get("runs") or []
with _lock:
conn = _connect()
try:
cur = conn.execute(
"INSERT INTO submissions(account_id,source_test_id,source_site,provider,model,test_name,"
"samples_ok,samples_total,avg_ttft_ms,avg_prefill_speed,avg_decode_speed,avg_stream_decode,"
"avg_prompt_tokens,avg_output_tokens,avg_total_ms,min_ttft_ms,max_ttft_ms,min_prefill_speed,"
"max_prefill_speed,min_decode_speed,max_decode_speed,min_total_ms,max_total_ms,"
"context_lengths,concurrency_levels,by_length,by_concurrency,runs,raw) VALUES("
"?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)",
(account_id,
int(data.get("source_test_id") or 0),
data.get("source_site") or "llm-speed-tester",
data.get("provider") or "",
data.get("model") or "",
data.get("test_name") or data.get("name") or "",
int(summary.get("samples_ok") or 0),
int(summary.get("samples_total") or 0),
summary.get("avg_ttft_ms"), summary.get("avg_prefill_speed"),
summary.get("avg_decode_speed"), summary.get("avg_stream_decode"),
summary.get("avg_prompt_tokens"), summary.get("avg_output_tokens"),
summary.get("avg_total_ms"),
summary.get("min_ttft_ms"), summary.get("max_ttft_ms"),
summary.get("min_prefill_speed"), summary.get("max_prefill_speed"),
summary.get("min_decode_speed"), summary.get("max_decode_speed"),
summary.get("min_total_ms"), summary.get("max_total_ms"),
json.dumps(gen.get("context_lengths") or summary.get("context_lengths") or [],
ensure_ascii=False),
json.dumps(summary.get("concurrency_levels") or gen.get("concurrency_levels") or [1],
ensure_ascii=False),
json.dumps(summary.get("by_length") or {}, ensure_ascii=False),
json.dumps(summary.get("by_concurrency") or {}, ensure_ascii=False),
json.dumps(runs[:200], ensure_ascii=False),
json.dumps(data, ensure_ascii=False)))
conn.commit()
return cur.lastrowid
finally:
conn.close()
def get_submission(sid: int):
with _lock:
conn = _connect()
try:
r = conn.execute("SELECT * FROM submissions WHERE id=?", (sid,)).fetchone()
if not r:
return None
d = dict(r)
d["account"] = conn.execute("SELECT name FROM accounts WHERE id=?",
(d["account_id"],)).fetchone()["name"]
for k in ("context_lengths", "concurrency_levels", "by_length",
"by_concurrency", "runs", "raw"):
d[k] = json.loads(d[k] or ("[]" if k in ("context_lengths", "concurrency_levels", "runs") else "{}"))
return d
finally:
conn.close()
def list_submissions(page=1, page_size=20, q="", account_id=None, model="", provider=""):
with _lock:
conn = _connect()
try:
where = ["1=1"]
args = []
if q:
where.append("(s.model LIKE ? OR s.test_name LIKE ? OR a.name LIKE ? OR s.provider LIKE ?)")
like = "%" + q + "%"
args += [like, like, like, like]
if account_id:
where.append("s.account_id=?")
args.append(account_id)
if model:
where.append("s.model=?")
args.append(model)
if provider:
where.append("s.provider=?")
args.append(provider)
wsql = " AND ".join(where)
total = conn.execute(
"SELECT COUNT(*) c FROM submissions s LEFT JOIN accounts a ON a.id=s.account_id WHERE " + wsql,
args).fetchone()["c"]
rows = conn.execute(
"SELECT s.*, a.name AS account FROM submissions s "
"LEFT JOIN accounts a ON a.id=s.account_id WHERE " + wsql +
" ORDER BY s.id DESC LIMIT ? OFFSET ?",
args + [page_size, (page - 1) * page_size]).fetchall()
out = []
for r in rows:
d = dict(r)
for k in ("context_lengths", "concurrency_levels"):
d[k] = json.loads(d[k] or "[]")
out.append(d)
return {"total": total, "items": out, "page": page, "page_size": page_size,
"pages": max(1, -(-total // page_size))}
finally:
conn.close()
def delete_submission(sid: int):
with _lock:
conn = _connect()
try:
conn.execute("DELETE FROM submissions WHERE id=?", (sid,))
conn.commit()
finally:
conn.close()
# ───────────────────────── 汇总统计 ─────────────────────────
def get_stats():
with _lock:
conn = _connect()
try:
model_cnt = conn.execute("SELECT COUNT(DISTINCT model) c FROM submissions WHERE model<>''").fetchone()["c"]
sub_cnt = conn.execute("SELECT COUNT(*) c FROM submissions").fetchone()["c"]
acc_cnt = conn.execute("SELECT COUNT(*) c FROM accounts").fetchone()["c"]
ok_cnt = conn.execute("SELECT SUM(samples_ok) c FROM submissions").fetchone()["c"] or 0
return {"models": model_cnt, "submissions": sub_cnt,
"accounts": acc_cnt, "samples_ok": ok_cnt}
finally:
conn.close()
def leaderboard(sort="avg_decode_speed", order="desc", limit=200):
"""按 (provider, model) 聚合所有提交 → 模型速度排行"""
sort_whitelist = {
"avg_decode_speed": "avg_decode_speed", "avg_prefill_speed": "avg_prefill_speed",
"avg_ttft_ms": "avg_ttft_ms", "best_decode": "best_decode",
"cnt": "cnt", "last_tested": "last_tested", "avg_stream_decode": "avg_stream_decode",
}
srt = sort_whitelist.get(sort, "avg_decode_speed")
od = "DESC" if order == "asc" and srt not in ("avg_ttft_ms", "avg_total_ms") else (
"ASC" if order == "asc" else "DESC")
with _lock:
conn = _connect()
try:
rows = conn.execute(
"SELECT provider, model, "
"COUNT(*) AS cnt, "
"COUNT(DISTINCT account_id) AS accounts, "
"AVG(avg_decode_speed) AS avg_decode_speed, "
"AVG(avg_prefill_speed) AS avg_prefill_speed, "
"AVG(avg_stream_decode) AS avg_stream_decode, "
"AVG(avg_ttft_ms) AS avg_ttft_ms, "
"AVG(avg_output_tokens) AS avg_output_tokens, "
"AVG(avg_total_ms) AS avg_total_ms, "
"MAX(avg_decode_speed) AS best_decode, "
"MAX(created_at) AS last_tested "
"FROM submissions WHERE model<>'' "
"GROUP BY provider, model ORDER BY %s %s LIMIT ?" % (srt, od),
(limit,)).fetchall()
return [dict(r) for r in rows]
finally:
conn.close()
def get_model_detail(provider, model):
"""单个模型的所有提交 + 按上下文长度聚合"""
with _lock:
conn = _connect()
try:
rows = conn.execute(
"SELECT s.*, a.name AS account FROM submissions s "
"LEFT JOIN accounts a ON a.id=s.account_id "
"WHERE s.model=? AND s.provider=? ORDER BY s.id DESC",
(model, provider)).fetchall()
subs = []
agg_length = {}
for r in rows:
d = dict(r)
d["by_length"] = json.loads(d.pop("by_length") or "{}")
d["by_concurrency"] = json.loads(d.pop("by_concurrency") or "{}")
d["context_lengths"] = json.loads(d.pop("context_lengths") or "[]")
d["concurrency_levels"] = json.loads(d.pop("concurrency_levels") or "[]")
subs.append(d)
for L, bl in d["by_length"].items():
a = agg_length.setdefault(str(L), {"count": 0, "sum_decode": 0.0,
"sum_prefill": 0.0, "sum_ttft": 0.0})
if bl.get("avg_decode_speed") is not None:
a["count"] += 1
a["sum_decode"] += float(bl["avg_decode_speed"] or 0)
a["sum_prefill"] += float(bl.get("avg_prefill_speed") or 0)
a["sum_ttft"] += float(bl.get("avg_ttft_ms") or 0)
length_rows = []
for L in sorted(agg_length, key=int):
a = agg_length[L]
if not a["count"]:
continue
length_rows.append({
"length": int(L), "count": a["count"],
"avg_decode_speed": round(a["sum_decode"] / a["count"], 2),
"avg_prefill_speed": round(a["sum_prefill"] / a["count"], 2),
"avg_ttft_ms": round(a["sum_ttft"] / a["count"], 1),
})
return {"provider": provider, "model": model,
"submissions": subs, "by_length": length_rows,
"sub_cnt": len(subs)}
finally:
conn.close()