Spaces:
Sleeping
Sleeping
Update inference.py
Browse files- inference.py +6 -6
inference.py
CHANGED
|
@@ -11,9 +11,9 @@ from typing import Dict, List, Optional, Tuple
|
|
| 11 |
from openai import OpenAI
|
| 12 |
|
| 13 |
try:
|
| 14 |
-
from clinical_trial_env import ClinicalTrialAction,
|
| 15 |
except ImportError:
|
| 16 |
-
from client import ClinicalTrialEnv
|
| 17 |
from models import ClinicalTrialAction
|
| 18 |
|
| 19 |
LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME") or os.getenv("IMAGE_NAME")
|
|
@@ -159,11 +159,11 @@ def format_action(action: ClinicalTrialAction) -> str:
|
|
| 159 |
return json.dumps(payload, separators=(",", ":"), sort_keys=True)
|
| 160 |
|
| 161 |
|
| 162 |
-
async def create_env() ->
|
| 163 |
if LOCAL_IMAGE_NAME:
|
| 164 |
-
return await
|
| 165 |
if ENV_BASE_URL:
|
| 166 |
-
env =
|
| 167 |
await env.connect()
|
| 168 |
return env
|
| 169 |
raise RuntimeError("Set LOCAL_IMAGE_NAME for Docker execution or ENV_BASE_URL for an existing server.")
|
|
@@ -171,7 +171,7 @@ async def create_env() -> ClinicalTrialEnv:
|
|
| 171 |
|
| 172 |
async def main() -> None:
|
| 173 |
client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)
|
| 174 |
-
env: Optional[
|
| 175 |
rewards: List[float] = []
|
| 176 |
steps_taken = 0
|
| 177 |
score = 0.5
|
|
|
|
| 11 |
from openai import OpenAI
|
| 12 |
|
| 13 |
try:
|
| 14 |
+
from clinical_trial_env import ClinicalTrialAction, ClinicalTrialEnvClient
|
| 15 |
except ImportError:
|
| 16 |
+
from client import ClinicalTrialEnv as ClinicalTrialEnvClient
|
| 17 |
from models import ClinicalTrialAction
|
| 18 |
|
| 19 |
LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME") or os.getenv("IMAGE_NAME")
|
|
|
|
| 159 |
return json.dumps(payload, separators=(",", ":"), sort_keys=True)
|
| 160 |
|
| 161 |
|
| 162 |
+
async def create_env() -> ClinicalTrialEnvClient:
|
| 163 |
if LOCAL_IMAGE_NAME:
|
| 164 |
+
return await ClinicalTrialEnvClient.from_docker_image(LOCAL_IMAGE_NAME)
|
| 165 |
if ENV_BASE_URL:
|
| 166 |
+
env = ClinicalTrialEnvClient(base_url=ENV_BASE_URL)
|
| 167 |
await env.connect()
|
| 168 |
return env
|
| 169 |
raise RuntimeError("Set LOCAL_IMAGE_NAME for Docker execution or ENV_BASE_URL for an existing server.")
|
|
|
|
| 171 |
|
| 172 |
async def main() -> None:
|
| 173 |
client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)
|
| 174 |
+
env: Optional[ClinicalTrialEnvClient] = None
|
| 175 |
rewards: List[float] = []
|
| 176 |
steps_taken = 0
|
| 177 |
score = 0.5
|