From f438206c0e72112c11d69fed8e41e54334c41d02 Mon Sep 17 00:00:00 2001 From: z3st Date: Fri, 24 Jul 2026 18:57:52 +0800 Subject: [PATCH] feat: submit only 2 models, add failed model tracking to avoid re-submission --- main.py | 41 +++++++++++++++++++++++++++++++++-------- 1 file changed, 33 insertions(+), 8 deletions(-) diff --git a/main.py b/main.py index bee1e13..bf5b240 100644 --- a/main.py +++ b/main.py @@ -82,16 +82,34 @@ db_conn = None def init_db(): global db_conn db_conn = sqlite3.connect(':memory:', check_same_thread=False) - db_conn.execute('''CREATE TABLE IF NOT EXISTS queue ( - model_id TEXT, gpu TEXT, url TEXT, downloads INTEGER, - params TEXT, category TEXT, score REAL, + db_conn.execute('''CREATE TABLE IF NOT EXISTS submitted ( + model_id TEXT, gpu TEXT, task_id TEXT, status TEXT, + submitted_at TEXT, checked_at TEXT, PRIMARY KEY(model_id, gpu) )''') - db_conn.execute('''CREATE TABLE IF NOT EXISTS submitted ( - model_id TEXT, gpu TEXT, task_id TEXT, submitted_at TEXT, + db_conn.execute('''CREATE TABLE IF NOT EXISTS failed ( + model_id TEXT, gpu TEXT, reason TEXT, failed_at TEXT, PRIMARY KEY(model_id, gpu) )''') db_conn.commit() +def is_model_failed(model_id: str) -> bool: + """检查模型是否已知失败""" + if db_conn: + row = db_conn.execute( + 'SELECT 1 FROM failed WHERE model_id=? AND gpu=?', + (model_id, TARGET_GPU) + ).fetchone() + return row is not None + return False + +def record_failed(model_id: str, reason: str): + """记录失败的模型""" + if db_conn: + db_conn.execute( + 'INSERT OR REPLACE INTO failed VALUES (?,?,?,?)', + (model_id, TARGET_GPU, reason, datetime.now().isoformat()) + ) + db_conn.commit() # ============================================================ @@ -302,7 +320,7 @@ def submit_model(model_url: str) -> tuple: # 主流程 # ============================================================ -def run_pipeline(submit_limit: int = 30): +def run_pipeline(submit_limit: int = 2): """完整流程:搜索→筛选→提交(只针对目标GPU)""" init_db() log("=" * 50) @@ -390,6 +408,10 @@ def run_pipeline(submit_limit: int = 30): if check_my_submitted(model_id): continue + # 检查是否已知失败(避免重复提交) + if is_model_failed(model_id): + continue + to_submit.append(m) if len(to_submit) >= min(submit_limit, available): break @@ -410,6 +432,9 @@ def run_pipeline(submit_limit: int = 30): ) else: log(f" ❌ {m['model_id']}: {msg}") + # 永久失败类型记录到 failed 表 + if any(kw in str(msg) for kw in ['保护期', '白名单', '唯一性']): + record_failed(m['model_id'], msg) time.sleep(0.5) log(f" 提交完成: {submitted}/{len(to_submit)}") @@ -470,7 +495,7 @@ class AgentHandler(BaseHTTPRequestHandler): if content_len > 0: body = json.loads(self.rfile.read(content_len)) - limit = body.get('limit', 30) + limit = body.get('limit', 2) self._json({'status': 'started', 'gpu': TARGET_GPU, 'limit': limit}) @@ -687,7 +712,7 @@ def main(): try: state['running'] = True state['last_run'] = datetime.now().isoformat() - count = run_pipeline(submit_limit=30) + count = run_pipeline(submit_limit=2) state['last_result'] = {'submitted': count, 'success': True} except Exception as e: log(f"流程异常: {traceback.format_exc()}")