SahilCodevally/codevally-vision-language-action
0
1"""2Language Reasoning Module — Natural-language command interpretation via LLM.3 4Takes a user command and the list of detected objects, then asks an LLM5to produce a structured action interpretation:6 7 {"action": "move", "source": "bottle", "destination": "box"}8 9Supported providers (set via LLM_PROVIDER env var):10 - "openai" → OpenAI API (GPT-4o-mini by default)11 - "groq" → Groq API (Llama 3.3 70B by default, ~5-10x faster)12"""13 14import json15import logging16from typing import Any17 18logger = logging.getLogger(__name__)19 20# ---------------------------------------------------------------------------21# System prompt sent to the LLM22# ---------------------------------------------------------------------------23_SYSTEM_PROMPT = """\24You are an AI assistant for a Vision Language Action robotics system.25 26Given:27 1. A list of objects detected in a scene (with labels and bounding boxes).28 2. A natural-language command from a human operator.29 30Your task:31 Interpret the command and return a JSON object with exactly these keys:32 - "action" : The verb describing what to do (e.g. "move", "pick_up",33 "place", "inspect", "rotate", "remove").34 - "source" : The object to act on.35 - "destination" : The target location or object (if applicable, else null).36 37CRITICAL RULES:38 • For "source" and "destination", you MUST use the EXACT label string from39 the detected objects list. The user may refer to objects using different40 words (e.g. "wrench" might actually be detected as "knife", "toolbox"41 might be "suitcase", "mug" might be "cup"). Your job is to figure out42 which detected object the user is referring to and use that detected label.43 • If the user says "wrench" but only "knife" is detected, return "knife".44 • If the user says "toolbox" but only "suitcase" is detected, return "suitcase".45 • If the user says "bin" but only "bowl" is detected, return "bowl".46 • NEVER return a label that is not in the detected objects list.47 • If truly no detected object can reasonably match, use the closest detected48 object label and note it.49 • Return ONLY the JSON object, no extra text.50 • If the command is ambiguous, make your best interpretation.51 • If no destination is implied, set "destination" to null.52"""53 54 55# Provider → base URL mapping (Groq uses the OpenAI-compatible SDK)56_PROVIDER_BASE_URLS: dict[str, str | None] = {57 "openai": None, # uses default OpenAI endpoint58 "groq": "https://api.groq.com/openai/v1",59}60 61 62class CommandReasoner:63 """64 Interprets natural-language commands using an OpenAI-compatible LLM.65 66 Supports OpenAI and Groq as backends — both expose the same67 Chat Completions API so the same prompt and parsing logic work for both.68 69 Attributes:70 provider: Active provider name ("openai" or "groq").71 model: Model identifier (e.g. "gpt-4o-mini" or "llama-3.3-70b-versatile").72 client: OpenAI SDK client instance (or None in fallback mode).73 """74 75 def __init__(76 self,77 provider: str,78 api_key: str,79 model: str,80 ):81 """82 Initialise the reasoner for a specific provider.83 84 Args:85 provider: "openai" or "groq".86 api_key: API key for the chosen provider.87 model: Model identifier string.88 """89 self.provider = provider90 self.model = model91 self.client = None92 93 if api_key:94 self._init_client(provider, api_key)95 else:96 logger.warning(97 "No API key provided for provider='%s' — "98 "CommandReasoner will use fallback mode.",99 provider,100 )101 102 def _init_client(self, provider: str, api_key: str) -> None:103 """Create an OpenAI-SDK client pointed at the chosen provider endpoint."""104 try:105 from openai import OpenAI106 107 base_url = _PROVIDER_BASE_URLS.get(provider) # None → OpenAI default108 kwargs: dict = {"api_key": api_key}109 if base_url:110 kwargs["base_url"] = base_url111 112 self.client = OpenAI(**kwargs)113 logger.info(114 "LLM client initialised (provider=%s, model=%s).",115 provider,116 self.model,117 )118 except Exception as exc:119 logger.error("Failed to initialise LLM client (%s): %s", provider, exc)120 self.client = None121 122 # ------------------------------------------------------------------123 # Public API124 # ------------------------------------------------------------------125 126 def interpret(127 self,128 command: str,129 detected_objects: list[dict[str, Any]],130 ) -> dict[str, Any]:131 """132 Interpret a user command in the context of detected objects.133 134 Args:135 command: Free-form user instruction.136 detected_objects: List of detection dicts from the vision module.137 138 Returns:139 Dict with keys ``action``, ``source``, ``destination``.140 """141 if not command or not command.strip():142 logger.warning("Empty command received.")143 return self._error_result("No command provided.")144 145 if self.client is None:146 logger.info("Using fallback reasoning (no LLM client).")147 return self._fallback_interpret(command, detected_objects)148 149 return self._llm_interpret(command, detected_objects)150 151 # ------------------------------------------------------------------152 # LLM-based interpretation153 # ------------------------------------------------------------------154 155 def _llm_interpret(156 self,157 command: str,158 detected_objects: list[dict[str, Any]],159 ) -> dict[str, Any]:160 """Call the LLM to interpret the command."""161 objects_description = self._format_objects(detected_objects)162 163 user_message = (164 f"Detected objects in the scene:\n{objects_description}\n\n"165 f"User command: \"{command}\"\n\n"166 f"Return the structured JSON interpretation."167 )168 169 try:170 response = self.client.chat.completions.create(171 model=self.model,172 messages=[173 {"role": "system", "content": _SYSTEM_PROMPT},174 {"role": "user", "content": user_message},175 ],176 temperature=0.1,177 max_tokens=256,178 )179 180 raw = response.choices[0].message.content.strip()181 logger.info("LLM raw response: %s", raw)182 183 return self._parse_response(raw)184 185 except Exception as exc:186 logger.error("LLM call failed: %s", exc)187 return self._fallback_interpret(command, detected_objects)188 189 # ------------------------------------------------------------------190 # Fallback (no LLM)191 # ------------------------------------------------------------------192 193 def _fallback_interpret(194 self,195 command: str,196 detected_objects: list[dict[str, Any]],197 ) -> dict[str, Any]:198 """199 Smart keyword-based fallback when no LLM is available.200 201 Uses the ObjectMapper for synonym-aware matching between user202 terms and detected YOLO labels.203 """204 from src.utils.object_mapper import find_best_match205 206 logger.info("Fallback interpretation for: '%s'", command)207 208 # Try to determine action verb209 action = "move"210 for verb in ("pick", "place", "move", "inspect", "rotate", "remove", "find"):211 if verb in command.lower():212 action = verb213 break214 215 labels = [obj["label"] for obj in detected_objects]216 cmd_lower = command.lower()217 218 # Extract noun phrases from command by removing common verbs/prepositions219 stop_words = {220 "the", "a", "an", "and", "or", "to", "from", "in", "on", "at",221 "into", "onto", "of", "with", "for", "it", "up", "back", "next",222 "pick", "place", "move", "inspect", "rotate", "remove", "find",223 "put", "take", "grab", "get", "bring", "carry", "push", "pull",224 "near", "close", "closest", "nearest", "under", "over", "above",225 "below", "beside", "behind",226 }227 words = cmd_lower.replace(",", " ").replace(".", " ").split()228 229 # Build candidate noun phrases (single words + bigrams)230 candidates: list[str] = []231 for w in words:232 if w not in stop_words and len(w) > 1:233 candidates.append(w)234 for i in range(len(words) - 1):235 bigram = f"{words[i]} {words[i+1]}"236 if words[i] not in stop_words or words[i+1] not in stop_words:237 candidates.append(bigram)238 239 # Match candidates against detected labels using ObjectMapper240 source = None241 destination = None242 used_labels: set[str] = set()243 244 for candidate in candidates:245 match = find_best_match(candidate, labels)246 if match and match not in used_labels:247 if source is None:248 source = match249 used_labels.add(match)250 elif destination is None:251 destination = match252 used_labels.add(match)253 break254 255 # Fallback: if no matches, use first two detected objects256 if source is None and labels:257 source = labels[0]258 if destination is None and len(labels) > 1:259 remaining = [l for l in labels if l != source]260 if remaining:261 destination = remaining[0]262 263 return {264 "action": action,265 "source": source or "unknown",266 "destination": destination,267 }268 269 # ------------------------------------------------------------------270 # Helpers271 # ------------------------------------------------------------------272 273 @staticmethod274 def _format_objects(detected_objects: list[dict[str, Any]]) -> str:275 """Format detected objects as a readable string for the LLM prompt."""276 if not detected_objects:277 return "No objects detected."278 279 lines: list[str] = []280 for i, obj in enumerate(detected_objects, 1):281 bbox = obj.get("bbox", [])282 conf = obj.get("confidence", 0)283 lines.append(284 f" {i}. {obj['label']} (confidence: {conf:.2f}, bbox: {bbox})"285 )286 return "\n".join(lines)287 288 @staticmethod289 def _parse_response(raw: str) -> dict[str, Any]:290 """291 Parse the LLM response string into a dict.292 293 Handles markdown code fences and extra whitespace.294 """295 cleaned = raw.strip()296 # Strip optional markdown code fences297 if cleaned.startswith("```"):298 lines = cleaned.split("\n")299 lines = [ln for ln in lines if not ln.strip().startswith("```")]300 cleaned = "\n".join(lines).strip()301 302 try:303 result = json.loads(cleaned)304 except json.JSONDecodeError:305 logger.warning("Failed to parse LLM JSON: %s", cleaned)306 return {307 "action": "unknown",308 "source": "unknown",309 "destination": None,310 "error": "Failed to parse LLM response.",311 "raw_response": raw,312 }313 314 # Ensure required keys315 return {316 "action": result.get("action", "unknown"),317 "source": result.get("source", "unknown"),318 "destination": result.get("destination"),319 }320 321 @staticmethod322 def _error_result(message: str) -> dict[str, Any]:323 """Return a structured error result."""324 return {325 "action": "error",326 "source": "unknown",327 "destination": None,328 "error": message,329 }330 