Spaces:
Sleeping
Sleeping
| """ | |
| HF Spaces server - exposes the AttentionEnv via a REST API | |
| compatible with the OpenEnv spec (reset / step / state endpoints). | |
| """ | |
| from fastapi import FastAPI, HTTPException | |
| from pydantic import BaseModel | |
| from typing import Dict, List, Optional | |
| import uvicorn | |
| from env.environment import AttentionEnv | |
| from env.models import Action | |
| from env.tasks import task_easy, task_medium, task_hard, grade_easy, grade_medium, grade_hard | |
| app = FastAPI(title="Attention Allocation System", version="1.0.0") | |
| _env: Optional[AttentionEnv] = None | |
| def get_env() -> AttentionEnv: | |
| global _env | |
| if _env is None: | |
| _env = AttentionEnv() | |
| return _env | |
| # ββ Request models βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| class StepRequest(BaseModel): | |
| item_id: int | |
| class BaselineResponse(BaseModel): | |
| scores: Dict[str, float] | |
| class GraderResponse(BaseModel): | |
| task: str | |
| score: float | |
| total_reward: float | |
| class TasksResponse(BaseModel): | |
| tasks: List[str] | |
| action_schema: Dict[str, str] | |
| # ββ Endpoints ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def root(): | |
| return { | |
| "name": "Attention Allocation System", | |
| "description": "Content recommendation RL environment (OpenEnv compatible)", | |
| "endpoints": ["/reset", "/step", "/state", "/health", "/baseline", "/grader", "/tasks"], | |
| } | |
| def health(): | |
| """Hackathon automated ping - must return 200.""" | |
| return {"status": "ok"} | |
| def reset(): | |
| """Start a new episode. Returns the initial observation.""" | |
| env = get_env() | |
| obs = env.reset() | |
| return {"observation": obs.model_dump(), "done": False} | |
| def step(request: StepRequest): | |
| """Take one action. Returns next observation, reward, and done flag.""" | |
| env = get_env() | |
| try: | |
| action = Action(item_id=request.item_id) | |
| obs, reward, done, info = env.step(action) | |
| return { | |
| "observation": obs.model_dump(), | |
| "reward": reward.value, | |
| "done": done, | |
| "info": info, | |
| } | |
| except Exception as e: | |
| raise HTTPException(status_code=400, detail=str(e)) | |
| def state(): | |
| """Get current observation without taking an action.""" | |
| env = get_env() | |
| if env.user is None: | |
| raise HTTPException(status_code=400, detail="Call /reset first") | |
| return env.state().model_dump() | |
| def baseline(): | |
| """ | |
| Run baseline inference on all 3 tasks (easy, medium, hard). | |
| Returns the baseline score for each task. | |
| """ | |
| from openai import OpenAI | |
| import os | |
| api_key = os.getenv("OPENAI_API_KEY") or os.getenv("HF_TOKEN") or os.getenv("API_KEY") | |
| if not api_key: | |
| raise HTTPException(status_code=500, detail="No API key configured for baseline inference") | |
| api_base_url = os.getenv("API_BASE_URL", "https://api.groq.com/openai/v1") | |
| model_name = os.getenv("MODEL_NAME", "llama-3.1-8b-instant") | |
| client = OpenAI(base_url=api_base_url, api_key=api_key) | |
| from inference import run_episode | |
| scores = {} | |
| tasks = [ | |
| ("easy", task_easy, 4.0), | |
| ("medium", task_medium, 7.0), | |
| ("hard", task_hard, 11.0), | |
| ] | |
| try: | |
| for task_name, task_fn, norm in tasks: | |
| env = task_fn() | |
| score = run_episode(client, env, task_name, norm) | |
| scores[task_name] = score | |
| return BaselineResponse(scores=scores) | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=f"Baseline inference failed: {str(e)}") | |
| def grader(): | |
| """ | |
| Calculate the grader score for the current episode. | |
| Returns the normalized score based on the number of items in the current environment. | |
| """ | |
| env = get_env() | |
| if env.user is None: | |
| raise HTTPException(status_code=400, detail="Call /reset first") | |
| if not hasattr(env, 'history') or len(env.history) == 0: | |
| raise HTTPException(status_code=400, detail="No actions taken yet in the episode") | |
| total_reward = getattr(env, 'total_reward', 0.0) | |
| if env.num_items == 5: | |
| task_name = "easy" | |
| score = grade_easy(total_reward) | |
| elif env.num_items == 10: | |
| task_name = "medium" | |
| score = grade_medium(total_reward) | |
| elif env.num_items == 15: | |
| task_name = "hard" | |
| score = grade_hard(total_reward) | |
| else: | |
| raise HTTPException(status_code=400, detail=f"Unknown task with {env.num_items} items") | |
| return GraderResponse(task=task_name, score=score, total_reward=total_reward) | |
| def tasks(): | |
| """ | |
| Returns the available tasks and the action schema required for each step. | |
| """ | |
| return TasksResponse( | |
| tasks=["easy", "medium", "hard"], | |
| action_schema={ | |
| "item_id": "integer (ID of the item to recommend)" | |
| } | |
| ) | |
| def main(): | |
| uvicorn.run(app, host="0.0.0.0", port=7860) | |
| if __name__ == "__main__": | |
| main() |