Team Ai
Apppublic

ruby56/Citation-Benchmark

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
environment.py147 linesDownload Raw Back to root
1import sqlite32import os3from typing import Optional, List, Dict, Any, Union, Tuple4from pydantic import BaseModel, Field5import tasks6 7class Observation(BaseModel):8    task_instruction: str9    search_results: Optional[List[Dict[str, str]]] = None10    last_abstract: Optional[str] = None11    citations_data: Optional[List[Dict[str, str]]] = None12    message: Optional[str] = None13    step_count: int = 014 15class Action(BaseModel):16    action_type: str = Field(..., description="One of: 'search', 'read_abstract', 'get_citations', 'submit_id'")17    query: Optional[str] = Field(None, description="Search term (used only for 'search')")18    paper_id: Optional[str] = Field(None, description="Paper ID (corpus_id or arxiv_id) to read, cite, or submit")19 20class Reward(BaseModel):21    value: float22    is_terminal: bool23    details: str24 25class CitationEnv:26    def __init__(self, task_id: str = "T001"):27        self.task = next((t for t in tasks.TASKS if t.id == task_id), tasks.TASKS[0])28        self.grader = tasks.Grader(self.task)29        self.step_count = 030        self.max_steps = 1531        self.state_history = []32        self._current_obs = None33        34        db_path = 'citation_db.sqlite'35        if not os.path.exists(db_path):36            print(f"Database {db_path} not found locally. Downloading from HF Datasets...")37            try:38                from huggingface_hub import hf_hub_download39                # Pulls directly from the Dataset instead of the Space40                db_path = hf_hub_download(repo_id="ruby56/Citation-Database", filename="citation_db.sqlite", repo_type="dataset" , local_dir=".")41                print("Database successfully mounted!")42            except Exception as e:43                print(f"WARNING: Could not download DB: {e}")44                45        self.db = sqlite3.connect(db_path)46        47    def reset(self) -> Observation:48        self.step_count = 049        self.state_history = []50        self._current_obs = Observation(51            task_instruction=self.task.claim,52            message=f"Environment initialized. Goal: {self.task.claim}"53        )54        self.state_history.append(self._current_obs)55        return self._current_obs56 57    def state(self) -> Observation:58        if not self._current_obs:59            return self.reset()60        return self._current_obs61 62    def step(self, action: Action) -> Tuple[Observation, Reward, bool, dict]:63        self.step_count += 164        done = False65        info = {"task": self.task.id}66        67        reward_value = -0.0568        reward_details = "Step taken."69 70        if self.step_count >= self.max_steps and action.action_type != "submit_id":71            return self._current_obs, Reward(value=-1.0, is_terminal=True, details="Max steps reached."), True, info72        73        message = ""74        search_results = self._current_obs.search_results if self._current_obs else None75        last_abstract = self._current_obs.last_abstract if self._current_obs else None76        citations_data = self._current_obs.citations_data if self._current_obs else None77 78        cursor = self.db.cursor()79 80        if action.action_type == "search":81            q = (action.query or "")82            query = "SELECT corpus_id, arxiv_id, title, year FROM s2_papers WHERE title LIKE ? LIMIT 10"83            cursor.execute(query, (f"%{q}%",))84            results = []85            for row in cursor.fetchall():86                results.append({"corpus_id": str(row[0]), "arxiv_id": str(row[1]), "title": row[2], "year": str(row[3])})87            search_results = results88            message = f"Found {len(results)} papers."89            90        elif action.action_type == "read_abstract":91            pid = action.paper_id92            cursor.execute("SELECT abstract FROM arxiv_metadata WHERE arxiv_id = ?", (pid,))93            res = cursor.fetchone()94            if res:95                last_abstract = res[0]96                message = f"Abstract for {pid} loaded."97            else:98                message = f"Abstract not found for {pid}. Make sure to pass a valid arxiv_id."99                reward_value -= 0.1100                101        elif action.action_type == "get_citations":102            pid = action.paper_id103            query = """104                SELECT c.contexts, c.intent, ar.title, ar.abstract105                FROM citations c106                JOIN s2_papers s2_citing ON c.citing_corpus_id = s2_citing.corpus_id107                JOIN arxiv_metadata ar ON s2_citing.arxiv_id = ar.arxiv_id108                WHERE c.cited_corpus_id = ?109                LIMIT 5110            """111            cursor.execute(query, (pid,))112            res = cursor.fetchall()113            if res:114                c_data = []115                for row in res:116                    c_data.append({"contexts": row[0], "intent": row[1], "citing_title": row[2], "citing_abstract": row[3]})117                citations_data = c_data118                message = f"Found {len(res)} valid citations with abstracts for corpus_id {pid}."119            else:120                message = f"No citation graph data found for corpus_id {pid}."121                122        elif action.action_type == "submit_id":123            pid = action.paper_id or ""124            score = self.grader.score(pid, self.db)125            reward_value = float(score)  # final reward is the grader score126            reward_details = f"Final evaluation. Score: {score}"127            message = f"ID submitted: {pid}"128            done = True129            130        else:131            message = "Invalid action."132            reward_value -= 0.2133 134        new_obs = Observation(135            task_instruction=self.task.claim,136            search_results=search_results,137            last_abstract=last_abstract,138            citations_data=citations_data,139            message=message,140            step_count=self.step_count141        )142        self._current_obs = new_obs143        self.state_history.append(new_obs)144        info["history_len"] = len(self.state_history)145 146        return new_obs, Reward(value=reward_value, is_terminal=done, details=reward_details), done, info147