nitdaa / agents /nvidia_llm.py
AI Agent
Update nvidia_llm to use minimaxai/minimax-m2.7 and new sampling params
fbc58fa
Raw
History Blame
6.85 kB
import os
import logging
import threading
import queue as _queue_module
import time
from flask import Flask, request, jsonify
from openai import OpenAI
# ── Logging ───────────────────────────────────────────────────────────────────
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(name)s] %(levelname)s %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
log = logging.getLogger("nvidia_llm")
logging.getLogger("werkzeug").setLevel(logging.ERROR)
logging.getLogger("httpx").setLevel(logging.WARNING)
# ── Config ────────────────────────────────────────────────────────────────────
HOST = os.getenv("NVIDIA_HOST", "127.0.0.1")
PORT = int(os.getenv("NVIDIA_PORT", "8002"))
NVIDIA_API_KEY = os.getenv("NVIDIA_API_KEY")
_QUEUE_MAX_SIZE = int(os.getenv("NVIDIA_QUEUE_MAX", "8"))
_REQUEST_TIMEOUT_S = int(os.getenv("NVIDIA_LLM_TIMEOUT", "600"))
DEFAULT_MODEL = "minimaxai/minimax-m2.7"
# ── Flask app ─────────────────────────────────────────────────────────────────
app = Flask(__name__)
@app.after_request
def after_request(response):
response.headers.add('Access-Control-Allow-Origin', '*')
response.headers.add('Access-Control-Allow-Headers', 'Content-Type,Authorization')
response.headers.add('Access-Control-Allow-Methods', 'GET,PUT,POST,DELETE,OPTIONS')
return response
# ── Inference Queue (multi-user serialization) ─────────────────────────────────
_inference_queue: _queue_module.Queue = _queue_module.Queue(maxsize=_QUEUE_MAX_SIZE)
if not NVIDIA_API_KEY:
log.warning("NVIDIA_API_KEY is not set. API calls might fail if the token is required.")
_nvidia_client = OpenAI(
base_url="https://integrate.api.nvidia.com/v1",
api_key=NVIDIA_API_KEY or "dummy-key-if-not-required-locally"
)
def _run_inference(data: dict) -> dict:
raw_prompt = data.get("prompt", "")
if not raw_prompt:
return {"error": "Field 'prompt' is required."}
prompts = raw_prompt if isinstance(raw_prompt, list) and (len(raw_prompt) == 0 or not isinstance(raw_prompt[0], dict)) else [raw_prompt]
max_tokens = int(data.get("max_tokens", 8192))
temperature = float(data.get("temperature", 1.0))
top_p = float(data.get("top_p", 0.95))
model_name = data.get("model", DEFAULT_MODEL)
choices = []
for i, prompt in enumerate(prompts):
if isinstance(prompt, list):
messages = prompt
else:
messages = [
{"role": "user", "content": prompt}
]
try:
print(f"\n[CONSOLE STREAM] Generating via NVIDIA for: {model_name}")
print("-" * 30)
completion = _nvidia_client.chat.completions.create(
model=model_name,
messages=messages,
temperature=temperature,
top_p=top_p,
max_tokens=max_tokens,
stream=False
)
full_output = ""
message = completion.choices[0].message
if message.content:
full_output = message.content
print(full_output)
print("\n" + "-" * 30)
choices.append({
"index": i,
"text": full_output.strip(),
"thinking": ""
})
except Exception as e:
log.error(f"Error calling NVIDIA API: {e}")
choices.append({
"index": i,
"text": f"Error: {str(e)}",
"thinking": ""
})
return {
"model": model_name,
"choices": choices,
"device": "cloud_nvidia"
}
def _inference_worker() -> None:
log.info("Inference worker thread started (pid=%d)", os.getpid())
while True:
try:
item = _inference_queue.get(timeout=1.0)
except _queue_module.Empty:
continue
req_data, result_holder, done_event = item
try:
result_holder[0] = _run_inference(req_data)
except Exception as exc:
log.error("Inference worker error: %s", exc)
result_holder[0] = {"error": f"Inference failed: {exc}"}
finally:
done_event.set()
_inference_queue.task_done()
_worker_thread = threading.Thread(target=_inference_worker, name="inference-worker", daemon=True)
_worker_thread.start()
# ── Routes ────────────────────────────────────────────────────────────────────
@app.route("/health", methods=["GET"])
def health():
return jsonify({
"status": "ok",
"queue_depth": _inference_queue.qsize(),
"queue_max": _QUEUE_MAX_SIZE
})
@app.route("/v1/completions", methods=["POST", "OPTIONS"])
def completions():
if request.method == "OPTIONS":
return jsonify({}), 200
data: dict = request.get_json(force=True) or {}
current_depth = _inference_queue.qsize()
if current_depth >= _QUEUE_MAX_SIZE:
return jsonify({
"error": "Server busy β€” all inference slots are occupied. Please try again shortly.",
"retry_after": 5
}), 503
result_holder: list = [None]
done_event = threading.Event()
try:
_inference_queue.put_nowait((data, result_holder, done_event))
except _queue_module.Full:
return jsonify({
"error": "Server busy β€” inference queue full. Please try again shortly.",
"retry_after": 5,
}), 503
completed = done_event.wait(timeout=_REQUEST_TIMEOUT_S)
if not completed:
return jsonify({
"error": f"Request timed out after {_REQUEST_TIMEOUT_S}s. ",
"retry_after": 10,
}), 503
result = result_holder[0]
if result is None:
return jsonify({"error": "Internal error: inference worker returned no result."}), 500
if "error" in result:
return jsonify(result), 500
return jsonify(result)
if __name__ == "__main__":
import signal, sys
def sigint_handler(sig, frame):
sys.exit(0)
signal.signal(signal.SIGINT, sigint_handler)
log.info(f"Starting NVIDIA LLM agent on http://{HOST}:{PORT}")
app.run(host=HOST, port=PORT, debug=False, threaded=True)