Team Ai
Apppublic

SahilCodevally/codevally-vision-language-action

sourceHugging Faceotherupdated 7mo agoView on Hugging Face
0likes
reasoner.py330 linesDownload Raw Back to language
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