Morget-01 / api /console.py
MGFeng's picture
Update api/console.py
55dd62f verified
Raw
History Blame Contribute Delete
7.61 kB
import time
from fastapi import APIRouter
from pydantic import BaseModel
import traceback
from config import MODEL_PRESETS, SORTED_SIZES, DATA_SOURCES, config
from core.state import workers, tasks, results, training_active, completed_task_ids
from core.task_manager import create_tasks, delete_completed_tasks, get_task_status
from core.database import save_results_to_hf
from core.state import training_active as global_training_active
from core.state import training_completed_tasks as global_training_completed_tasks
router = APIRouter(prefix="/api/console", tags=["console"])
class ConsoleCommand(BaseModel):
command: str
@router.post("/command")
async def execute_command(cmd: ConsoleCommand):
try:
c = cmd.command.strip()
if not c:
return {"output": "Please enter a command"}
parts = c.split()
cmd_name = parts[0].lower()
args = parts[1:]
if cmd_name == "help":
return {"output": f"""
Available commands:
status - Show status
workers - Show worker list
tasks - Show task status
results - Show training results
create <size> - Create training tasks ({len(MODEL_PRESETS)} presets available)
create <size> <datasource> - Specify data source
presets - Show all presets
datasources - Show all data sources
delete-completed - Delete all completed tasks
save-results - Force save results to Hugging Face
stop - Stop training
clear - Clear all tasks
config - Show config
set <k>=<v> - Set config value
help - Show this help
"""}
elif cmd_name == "datasources":
return {"output": "πŸ“Š Available data sources:\n" + "\n".join([f" {i+1}. {s}" for i, s in enumerate(DATA_SOURCES)])}
elif cmd_name == "presets":
output = f"πŸ“Š Available model presets ({len(MODEL_PRESETS)} total):\n"
for size, info in MODEL_PRESETS.items():
output += f" {size}: tasks={info.get('tasks', 0)}\n"
return {"output": output}
elif cmd_name == "status":
s = get_task_status()
online = sum(1 for w in workers.values() if time.time() - w.get("last_heartbeat", 0) < 90)
return {"output": f"""
πŸ“Š Status:
Training: {'🟒 Active' if s['training_active'] else 'πŸ”΄ Stopped'}
Target: {s['target']}
Tasks: {s['completed_tasks']}/{s['total_tasks']} ({s['progress']:.1f}%)
Workers: {online}/{len(workers)} online
Results: {len(results)} in memory
Total completed: {s['total_completed']}
Data source: {config.get('data_source')}
"""}
elif cmd_name == "workers":
from api.workers import list_workers
wl = await list_workers()
output = "πŸ–₯️ Worker list:\n"
for w in wl["workers"]:
status_icon = "🟒" if w["status"] == "online" else "πŸ”΄"
output += f" {status_icon} {w['worker_id']}: {w['status_detail']} | progress: {w['progress']*100:.0f}% | completed: {w['completed_tasks']}\n"
if w.get("current_task"):
output += f" current: {w['current_task']}\n"
return {"output": output}
elif cmd_name == "tasks":
s = get_task_status()
return {"output": f"""
πŸ“‹ Task status:
Pending: {s['pending']}
Assigned: {s['assigned']}
Completed: {s['completed']}
Total: {s['total']}
Progress: {s['progress']:.1f}%
Total completed: {s['total_completed']}
"""}
elif cmd_name == "results":
from api.results import get_results
r = await get_results()
output = f"πŸ“ˆ Training results ({r['total']} total):\n"
for res in r["results"][-20:]:
output += f" {res['task_id']}: loss={res.get('loss', '?')}, worker={res.get('worker_id', '?')[:12]}\n"
return {"output": output}
elif cmd_name == "create":
if len(args) < 1:
return {"output": "❌ Usage: create <size> [datasource]"}
target_input = args[0]
data_source = args[1] if len(args) >= 2 else DATA_SOURCES[0]
language = args[2] if len(args) >= 3 else "python"
matched_target = None
for key in MODEL_PRESETS.keys():
if key.lower() == target_input.lower():
matched_target = key
break
if matched_target is None:
available = ', '.join(list(MODEL_PRESETS.keys())[:15])
return {"output": f"❌ Unknown target: {target_input}\nAvailable: {available}... (total {len(MODEL_PRESETS)})"}
if data_source not in DATA_SOURCES:
return {"output": f"❌ Unknown data source: {data_source}\nAvailable: {DATA_SOURCES}"}
result = create_tasks(matched_target, language, data_source)
return {"output": f"""
βœ… Training tasks created:
Target: {result['target_size']}
Language: {result['language']}
Data source: {result['data_source']}
Number of tasks: {result['num_tasks']}
Steps per task: {result['steps_per_task']}
Total steps: {result['total_steps']}
"""}
elif cmd_name == "delete-completed":
count = delete_completed_tasks()
return {"output": f"βœ… Deleted {count} completed tasks"}
elif cmd_name == "save-results":
success = save_results_to_hf(results)
return {"output": f"{'βœ…' if success else '❌'} Saved {len(results)} results"}
elif cmd_name == "stop":
global_training_active = False
save_results_to_hf(results)
return {"output": "⏹️ Training stopped, results saved"}
elif cmd_name == "clear":
tasks.clear()
results.clear()
global_training_active = False
global_training_completed_tasks = 0
completed_task_ids.clear()
for wid in workers:
workers[wid]["current_task"] = None
workers[wid]["backup_task"] = None
workers[wid]["status"] = "idle"
return {"output": "πŸ—‘οΈ Cleared all tasks"}
elif cmd_name == "config":
current = config.get_all()
output = "βš™οΈ Current config:\n"
for k, v in current.items():
output += f" {k}: {v}\n"
return {"output": output}
elif cmd_name == "set" and len(args) >= 1:
try:
key, value = args[0].split("=")
if key in config.get_all():
old = config.get(key)
if isinstance(old, bool):
config.set(key, value.lower() in ("true", "1", "yes"))
elif isinstance(old, int):
config.set(key, int(float(value)))
elif isinstance(old, float):
config.set(key, float(value))
else:
config.set(key, value)
return {"output": f"βœ… {key} = {config.get(key)}"}
else:
return {"output": f"❌ Unknown config key: {key}"}
except Exception as e:
return {"output": f"❌ Format error: {e}"}
else:
return {"output": f"❌ Unknown command: {cmd_name}\nType help for available commands"}
except Exception as e:
print(f"❌ Command error: {e}")
print(traceback.format_exc())
return {"output": f"❌ Command execution error: {str(e)}"}