Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
65 changes: 65 additions & 0 deletions nemo_gym/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
import asyncio
import json
import shlex
from copy import deepcopy
from glob import glob
from os import environ, makedirs
from os.path import exists
Expand Down Expand Up @@ -89,6 +90,70 @@ class RunHelper: # pragma: no cover
_processes: Dict[str, Popen]
_server_instances: List[ServerInstance]

def parse_from_dynamic_or_defaults(
self,
initial_global_config: dict,
head_server_host: Optional[str] = None,
head_server_port: Optional[int] = None,
policy_model_name: Optional[str] = None,
policy_base_url: Optional[str] = None,
policy_api_key: Optional[str] = None,
) -> dict:
dynamic_cfg = initial_global_config.get("dynamic", None)
if not dynamic_cfg:
print(f"DEBUG: RunHelper: warning: no dynamic config, using provided defaults...", flush=True)
initial_global_config["head_server"] = {
"host": head_server_host,
"port": head_server_port,
}
initial_global_config["policy_model_name"] = policy_model_name
initial_global_config["policy_base_url"] = policy_base_url
initial_global_config["policy_api_key"] = policy_api_key
return initial_global_config

print(f"DEBUG: RunHelper: using dynamic config...", flush=True)

# Special top-level config keys (head server, policy).
initial_global_config["head_server"] = {
"host": dynamic_cfg["head_server"]["host"],
"port": dynamic_cfg["head_server"]["port"],
}
initial_global_config["policy_model_name"] = dynamic_cfg["policy_model"]["model_name"]
initial_global_config["policy_base_url"] = dynamic_cfg["policy_model"]["base_url"]
initial_global_config["policy_api_key"] = "dummy_key"

def _merge_dict_inplace(target: dict, source: dict):
for src_key, src_value in source.items():
if (
src_key in target and
isinstance(target[src_key], dict) and
isinstance(src_value, dict)
):
_merge_dict_inplace(target[src_key], src_value)
else:
target[src_key] = deepcopy(src_value)

# Merge all other top-level config keys.
for key, sub_cfg in dynamic_cfg.items():
if key in ("head_server", "policy_model", ):
continue
elif key in NEMO_GYM_RESERVED_TOP_LEVEL_KEYS:
raise ValueError(
f"found reserved top-level key {repr(key)} in 'dynamic' config section"
)
elif key not in initial_global_config:
print(f"RunHelper: warning: dynamic top-level key {repr(key)} is not present in the global config, skipping...", flush=True)
continue
elif initial_global_config[key] is None:
initial_global_config[key] = {}
elif not isinstance(initial_global_config[key], dict):
raise ValueError(
f"expected top-level key {repr(key)} to be a dict, got {type(initial_global_config[key]).__name__}"
)
_merge_dict_inplace(initial_global_config[key], sub_cfg)

return initial_global_config

def start(self, global_config_dict_parser_config: GlobalConfigDictParserConfig) -> None:
global_config_dict = get_global_config_dict(global_config_dict_parser_config=global_config_dict_parser_config)

Expand Down
2 changes: 2 additions & 0 deletions nemo_gym/global_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,11 +35,13 @@
ENTRYPOINT_KEY_NAME = "entrypoint"
DEFAULT_HOST_KEY_NAME = "default_host"
HEAD_SERVER_KEY_NAME = "head_server"
DYNAMIC_KEY_NAME = "dynamic"
NEMO_GYM_RESERVED_TOP_LEVEL_KEYS = [
CONFIG_PATHS_KEY_NAME,
ENTRYPOINT_KEY_NAME,
DEFAULT_HOST_KEY_NAME,
HEAD_SERVER_KEY_NAME,
DYNAMIC_KEY_NAME,
]

POLICY_BASE_URL_KEY_NAME = "policy_base_url"
Expand Down
12 changes: 10 additions & 2 deletions nemo_gym/server_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,9 +54,11 @@
# Eventually, we may also want to parameterize the max connections. For now, we set the max connections to just some very large number.
#
# It's critical that this client is NOT used before uvicorn.run is called. Under the hood, this async client will start and use an event loop, and store a handle to that specific event loop. When uvicorn.run is called, it will replace the event loop policy with its own. So the handle that the async client has is now outdated.
_MAX_CONNECTIONS = getenv("NEMO_GYM_HTTPX_MAX_CONNECTIONS", 9001)
_MAX_RETRIES = getenv("NEMO_GYM_HTTPX_MAX_RETRIES", 3)
GLOBAL_HTTPX_CLIENT = AsyncClient(
limits=Limits(max_keepalive_connections=1500, max_connections=1500),
transport=AsyncHTTPTransport(retries=3),
limits=Limits(max_keepalive_connections=_MAX_CONNECTIONS, max_connections=_MAX_CONNECTIONS),
transport=AsyncHTTPTransport(retries=_MAX_RETRIES),
timeout=None,
)

Expand All @@ -81,21 +83,27 @@ def load_head_server_config(cls) -> BaseServerConfig:
@classmethod
def load_from_global_config(cls, head_server_config: Optional[BaseServerConfig] = None) -> "ServerClient":
if head_server_config is None:
print(f"DEBUG: ServerClient.load_from_global_config: head_server_config is None, load...", flush=True)
head_server_config = cls.load_head_server_config()
print(f"DEBUG: ServerClient.load_from_global_config: head_server_config = {head_server_config}", flush=True)

# It's critical we use requests here instead of the global httpx client since a FastAPI server may be run downstream of this function call.
head_server_url = f"http://{head_server_config.host}:{head_server_config.port}"
print(f"DEBUG: ServerClient.load_from_global_config: connect: head_server_url = {head_server_url}", flush=True)
try:
response = requests.get(
f"{head_server_url}/global_config_dict_yaml",
)
print(f"DEBUG: ServerClient.load_from_global_config: connect: response = {response}", flush=True)
except ConnectionError as e:
print(f"DEBUG: ServerClient.load_from_global_config: connect: except = {e}", flush=True)
raise ValueError(
f"Could not connect to the head server at {head_server_url}. Perhaps you are not running a server or your head server is on a different port?"
) from e

global_config_dict_yaml = response.content.decode()
global_config_dict = OmegaConf.create(json.loads(global_config_dict_yaml))
print(f"DEBUG: ServerClient.load_from_global_config: global_config_dict = {global_config_dict}", flush=True)

return cls(head_server_config=head_server_config, global_config_dict=global_config_dict)

Expand Down
106 changes: 104 additions & 2 deletions resources_servers/library_judge_math/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import asyncio
import contextlib
import logging
from io import StringIO
Expand All @@ -21,6 +22,7 @@
from math_verify.errors import TimeoutException
from math_verify.metric import math_metric
from math_verify.parser import ExprExtractionConfig, LatexExtractionConfig
from openai import OpenAI
from pydantic import BaseModel

from nemo_gym.base_resources_server import (
Expand All @@ -36,6 +38,9 @@
NeMoGymResponse,
NeMoGymResponseCreateParamsNonStreaming,
)
from nemo_gym.server_utils import (
get_global_config_dict,
)


class LibraryJudgeMathResourcesServerConfig(BaseResourcesServerConfig):
Expand Down Expand Up @@ -107,6 +112,86 @@ def model_post_init(self, context: Any) -> None:
),
)

global_cfg = get_global_config_dict()
dynamic_cfg = global_cfg.get("dynamic", dict())
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: global cfg = {global_cfg}", flush=True)
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: dynamic cfg = {dynamic_cfg}", flush=True)

if False:
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: judge init: ...", flush=True)
judge_model_name = None
judge_base_url = None
judge_client = None
self._judge_model_name = None
self._judge_client = None
for _, judge in dynamic_cfg.get("judges", dict()).items():
# TODO(peter): select judge by provided "capability".
judge_model_name = judge["model_name"]
judge_base_url = judge["generation_base_url"]
judge_client = OpenAI(
base_url=judge_base_url,
api_key="dummy_key",
)
# judge_models = judge_client.models.list()
self._judge_model_name = judge_model_name
self._judge_client = judge_client
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: judge init: ok", flush=True)
break

if False:
# if self._judge_client is not None:
test_request = {
"model": judge_model_name,
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 512,
"temperature": 0.6,
"top_p": 1.0,
}
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: /v1/chat/completions test request = {test_request}", flush=True)
test_response = judge_client.chat.completions.create(
**test_request
)
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: /v1/chat/completions test response = {test_response}", flush=True)

test_request = {
"model": judge_model_name,
"input": [{"role": "user", "content": "hi"}],
"max_output_tokens": 512,
"temperature": 0.6,
"top_p": 1.0,
}
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: /v1/responses test request = {test_request}", flush=True)
try:
test_response = asyncio.run(self.server_client.post(
server_name="math_judge",
url_path="/v1/responses",
json=test_request,
))
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: /v1/responses test response = {test_response}", flush=True)
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: /v1/responses test response payload = {test_response.json()}", flush=True)
except Exception as e:
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: /v1/responses test exception = {e}", flush=True)

if False:
test_request = {
"model": "deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B",
"input": [{"role": "user", "content": "hi"}],
"max_output_tokens": 512,
"temperature": 0.6,
"top_p": 1.0,
}
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: /v1/responses test request 2 = {test_request}", flush=True)
try:
test_response = asyncio.run(self.server_client.post(
server_name="math_judge",
url_path="/v1/responses",
json=test_request,
))
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: /v1/responses test response 2 = {test_response}", flush=True)
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: /v1/responses test response 2 payload = {test_response.json()}", flush=True)
except Exception as e:
print(f"DEBUG: LibraryJudgeMathResourcesServer.model_post_init: /v1/responses test exception = {e}", flush=True)

def setup_webserver(self) -> FastAPI:
app = super().setup_webserver()

Expand Down Expand Up @@ -247,12 +332,29 @@ async def _generate_judge_evaluation(
),
]

judge_name = config.judge_model_server.name
# print(f"DEBUG: LibraryJudgeMathResourcesServer._generate_judge_evaluation: judge name = {repr(judge_name)} payload = {responses_create_params}", flush=True)
response = await self.server_client.post(
server_name=config.judge_model_server.name,
server_name=judge_name,
url_path="/v1/responses",
json=responses_create_params,
)
judge_response = NeMoGymResponse.model_validate(response.json())
try:
response = await self.server_client.post(
server_name=judge_name,
url_path="/v1/responses",
json=responses_create_params,
)
except Exception as e:
print(f"DEBUG: LibraryJudgeMathResourcesServer._generate_judge_evaluation: except = {e} server = {repr(judge_name)} responses create params = {responses_create_params}", flush=True)
judge_evaluation = JudgeEvaluation(
responses_create_params=responses_create_params,
response={},
)
return False, judge_evaluation
response_body = response.json()
# print(f"DEBUG: LibraryJudgeMathResourcesServer._generate_judge_evaluation: response = {response} payload = {response_body}", flush=True)
judge_response = NeMoGymResponse.model_validate(response_body)
judge_evaluation = JudgeEvaluation(responses_create_params=responses_create_params, response=judge_response)

# Currently, for all the cases in which the response from the LLM judge
Expand Down
11 changes: 10 additions & 1 deletion responses_api_agents/simple_agent/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,7 +121,16 @@ async def run(self, body: SimpleAgentRunRequest) -> SimpleAgentVerifyResponse:
url_path="/verify",
json=verify_request.model_dump(),
)
return SimpleAgentVerifyResponse.model_validate(verify_response.json())
try:
verify_response_body = verify_response.json()
except Exception as e:
print(f"DEBUG: SimpleAgent.run: except = {e} type(response) = {type(verify_response).__name__} response = {verify_response}", flush=True)
try:
print(f"DEBUG: SimpleAgent.run: model dump = {verify_response.model_dump()}", flush=True)
except Exception as e2:
print(f"DEBUG: SimpleAgent.run: no model dump: {e2}", flush=True)
raise e
return SimpleAgentVerifyResponse.model_validate(verify_response_body)


if __name__ == "__main__":
Expand Down
10 changes: 10 additions & 0 deletions responses_api_models/vllm_model/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,13 @@ class VLLMModel(SimpleResponsesAPIModel):
config: VLLMModelConfig

def model_post_init(self, context):
if False:
port = self.config.port
base_url = self.config.base_url
print(f"DEBUG: VLLMModel.model_post_init: port = {port} ctx = {context} base_url = {repr(base_url)}", flush=True)
if not base_url.endswith("/v1"):
base_url = f"{base_url}/v1"
print(f"DEBUG: VLLMModel.model_post_init: port = {port} ctx = {context} base_url = {repr(base_url)} (with /v1)", flush=True)
self._client = NeMoGymAsyncOpenAI(
base_url=self.config.base_url,
api_key=self.config.api_key,
Expand All @@ -77,8 +84,11 @@ def model_post_init(self, context):
return super().model_post_init(context)

async def responses(self, body: NeMoGymResponseCreateParamsNonStreaming = Body()) -> NeMoGymResponse:
# print(f"DEBUG: VLLMModel.responses: body = {body}", flush=True)

# Response Create Params -> Chat Completion Create Params
chat_completion_create_params = self._converter.responses_to_chat_completion_create_params(body)
# print(f"DEBUG: VLLMModel.responses: chat completion payload = {chat_completion_create_params}", flush=True)
if not body.model:
body.model = self.config.model

Expand Down
Loading