Anoopsingh53 commited on
Commit
4a7772e
·
verified ·
1 Parent(s): c12e5b3

Upload inference.py

Browse files
Files changed (1) hide show
  1. inference.py +6 -5
inference.py CHANGED
@@ -6,11 +6,12 @@ from models import RiderSafetyAction, RiderSafetyObservation
6
  from client import RiderSafetyClient
7
 
8
  # 1. Configuration
9
- API_BASE_URL = os.getenv("API_BASE_URL", "https://api-inference.huggingface.co/v1/")
10
- MODEL_NAME = os.getenv("MODEL_NAME", "meta-llama/Meta-Llama-3-8B-Instruct")
11
- HF_TOKEN = os.getenv("HF_TOKEN")
 
12
 
13
- client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)
14
 
15
  # 2. Logging Helpers (Standard Format)
16
  def log_start(task, env_name, model):
@@ -93,7 +94,7 @@ async def run_task(env: RiderSafetyClient, task_name: str):
93
  log_end(success_flag, steps, score, rewards)
94
 
95
  async def main():
96
- env = RiderSafetyClient("http://localhost:8000")
97
  for task in ["easy", "medium", "hard"]:
98
  await run_task(env, task)
99
  print("-" * 50)
 
6
  from client import RiderSafetyClient
7
 
8
  # 1. Configuration
9
+ # The validator requires using API_BASE_URL and API_KEY environment variables exactly.
10
+ API_BASE_URL = os.environ.get("API_BASE_URL")
11
+ API_KEY = os.environ.get("API_KEY")
12
+ MODEL_NAME = os.environ.get("MODEL_NAME", "meta-llama/Meta-Llama-3-8B-Instruct")
13
 
14
+ client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)
15
 
16
  # 2. Logging Helpers (Standard Format)
17
  def log_start(task, env_name, model):
 
94
  log_end(success_flag, steps, score, rewards)
95
 
96
  async def main():
97
+ env = RiderSafetyClient("http://localhost:7860")
98
  for task in ["easy", "medium", "hard"]:
99
  await run_task(env, task)
100
  print("-" * 50)