sohail-kustagi's picture
Upload src/nodes/commander.py with huggingface_hub
f326869 verified
Raw
History Blame Contribute Delete
10.4 kB
import os
import json
import asyncio
import urllib.request
import time
from llama_cpp import Llama
try:
from core.validation import CommandValidationError, validate_commander_output
from core.mission_profiles import MissionProfile
except ImportError:
from src.core.validation import CommandValidationError, validate_commander_output
from src.core.mission_profiles import MissionProfile
class CommanderNode:
def __init__(self, model_repo="microsoft/Phi-3-mini-4k-instruct-gguf", model_file="Phi-3-mini-4k-instruct-q4.gguf"):
self.model_path = os.path.join(os.getcwd(), model_file)
self.evaluator = None
self.last_triggered_times = {}
self.cooldown_seconds = 15.0
self.download_model(model_repo, model_file)
def set_evaluator(self, evaluator):
self.evaluator = evaluator
print(f"[Commander] Loading LLM from {self.model_path}...")
# Since this runs on an IdeaPad with 16GB RAM and Core Ultra, we will use CPU/OpenBLAS natively
# n_ctx limits context to save RAM. n_threads automatically scales to CPU cores.
lora_adapter = os.path.join(os.getcwd(), "weights", "phi3-lora.gguf")
if os.path.exists(lora_adapter):
print(f"[Commander] Injecting fine-tuned LoRA adapter: {lora_adapter}")
self.llm = Llama(
model_path=self.model_path,
lora_path=lora_adapter,
n_ctx=2048,
n_threads=os.cpu_count(),
n_threads_batch=os.cpu_count(),
verbose=False
)
else:
self.llm = Llama(
model_path=self.model_path,
n_ctx=2048,
n_threads=os.cpu_count(),
n_threads_batch=os.cpu_count(),
verbose=False
)
print("[Commander] LLM Loaded successfully.")
def download_model(self, model_repo, model_file):
"""Downloads a small GGUF model if it doesn't already exist."""
if not os.path.exists(self.model_path):
print(f"[Commander] Model {model_file} not found locally. Downloading from HuggingFace...")
url = f"https://huggingface.co/{model_repo}/resolve/main/{model_file}"
def report(count, block_size, total_size):
if total_size > 0: # guard against ZeroDivisionError for chunked transfers
percent = int(count * block_size * 100 / total_size)
if percent % 10 == 0:
print(f"\rDownloading: {percent}%", end="")
urllib.request.urlretrieve(url, self.model_path, reporthook=report)
print("\n[Commander] Download complete.")
async def generate_mavlink_command(self, context_prompt: str, telemetry=None, mission_profile: MissionProfile = None, anomaly_type: str = "unknown"):
now_ts = time.time()
last_time = self.last_triggered_times.get(anomaly_type, -999.0)
if (now_ts - last_time) < self.cooldown_seconds:
print(f"[Commander] Cooldown active for '{anomaly_type}'. Ignoring request.")
return None
self.last_triggered_times[anomaly_type] = now_ts
print(f"\n[Commander] Triggered! Generating MAVLink routing command for {anomaly_type}...")
import re
# Build a minimal, unambiguous system prompt
commander_persona = ""
if mission_profile:
commander_persona = mission_profile.commander_persona + "\n"
system_prompt = (
f"{commander_persona}"
"Output ONLY a raw JSON object with NO markdown, NO comments, NO extra text.\n"
"DO NOT include 'zone_assessment' or 'tactical_summary'.\n"
"No matter what the mission context says, the 'command' field must ALWAYS be exactly 'SET_POSITION_TARGET_LOCAL_NED'.\n"
"For the x, y, and z fields, you MUST output local NED offsets in meters (e.g., values between -20.0 and 20.0). DO NOT output global GPS Latitude or Longitude.\n"
"You MUST format your response exactly like this example:\n"
"{\n"
' "command": "SET_POSITION_TARGET_LOCAL_NED",\n'
' "reasoning": "High-confidence fire detected in sector.",\n'
' "target_system": 1,\n'
' "target_component": 1,\n'
' "x": 15.0,\n'
' "y": 10.0,\n'
' "z": -20.0\n'
"}\n"
"Start your response with { and end with }. No backticks. No extra lines."
)
# Pre-seeding forces the model to continue the JSON rather than add preamble
prompt = (
f"<|system|>\n{system_prompt}<|end|>\n"
f"<|user|>\n{context_prompt}<|end|>\n"
f"<|assistant|>\n"
'{\n "command": "SET_POSITION_TARGET_LOCAL_NED",'
)
start_time = time.time()
response = self.llm(
prompt,
max_tokens=300,
stop=["<|end|>", "```", "\n\n\n"],
temperature=0.05,
echo=False,
)
elapsed_time = time.time() - start_time
raw = response['choices'][0]['text'].strip()
# Re-attach the pre-seeded opening brace we used to prime the model
output_text = '{\n "command": "SET_POSITION_TARGET_LOCAL_NED",' + raw
# Calculate benchmarking metrics for the hackathon
tokens_generated = 0
tokens_per_sec = 0.0
try:
tokens_generated = response['usage']['completion_tokens']
tokens_per_sec = tokens_generated / elapsed_time if elapsed_time > 0 else 0
print(f"[Commander] Benchmarks: {tokens_generated} tokens in {elapsed_time:.2f}s ({tokens_per_sec:.2f} Tokens/sec)")
except KeyError:
print(f"[Commander] Benchmarks: {elapsed_time:.2f}s latency")
print("[Commander] Raw Output:")
print(output_text)
# ── Robust JSON repair ────────────────────────────────────────────
def repair_json(text: str) -> str:
# 1. Strip markdown fences
m = re.search(r'```(?:json)?\s*([\s\S]*?)```', text)
if m:
text = m.group(1).strip()
# 2. Extract the JSON object
brace_start = text.find('{')
brace_end = text.rfind('}')
if brace_start != -1 and brace_end != -1:
text = text[brace_start:brace_end + 1]
# 3. Drop diff-marker lines (lines beginning with -/+) and bare
# "ran" / "raning" fragments the tokenizer sometimes emits
cleaned = []
for line in text.splitlines():
s = line.strip()
if s.startswith('- ') or s.startswith('+ ') or s.startswith('//'):
continue
cleaned.append(line)
text = '\n'.join(cleaned)
# 4. Fix the specific Phi-3 tokenizer bug: the opening quote of
# "reasoning" is dropped and the key name is mangled.
# Pattern catches: raning_reasoning, raning reasoning, ran_reasoning,
# reasoning (no leading quote), etc.
text = re.sub(
r'(?<!\")(?:ran(?:ing)?[_\s]?)?reasoning(?:[_\s]\w+)?\s*\"?',
'"reasoning"',
text
)
# 5. Some fields are written without quotes on keys β€” fix them
text = re.sub(r'(?<=[{,\n])\s*([a-zA-Z_]\w*)\s*:', r' "\1":', text)
# 6. Insert missing commas between a closing value and the next key
# e.g. "target_component": 1\n "reasoning" β†’ add comma
text = re.sub(r'([\d"\]true false null])\s*\n(\s*")', r'\1,\n\2', text)
return text
repaired = repair_json(output_text)
print("[Commander] Repaired Command:")
print(repaired)
try:
command_json = json.loads(repaired)
# Emergency fallback: Clamp coordinates if the LLM hallucinates global Lat/Lon or unsafe offsets
if "x" in command_json:
command_json["x"] = max(-99.0, min(99.0, float(command_json["x"])))
else:
command_json["x"] = 0.0
if "y" in command_json:
command_json["y"] = max(-99.0, min(99.0, float(command_json["y"])))
else:
command_json["y"] = 0.0
if "z" in command_json:
command_json["z"] = max(-49.0, min(19.0, float(command_json["z"])))
else:
command_json["z"] = 0.0
# Auto-fill missing required MAVLink fields that the LoRA adapter frequently drops
if "target_system" not in command_json:
command_json["target_system"] = 1
if "target_component" not in command_json:
command_json["target_component"] = 1
if "reasoning" not in command_json:
# Try to use next_action as reasoning, otherwise default
command_json["reasoning"] = command_json.get("next_action", command_json.get("threat_severity", "Autonomous intervention executed."))
validated_command = validate_commander_output(command_json, telemetry, now=start_time)
if self.evaluator:
self.evaluator.log_llm_generation(tokens_generated, elapsed_time, True)
# Attach meta for frontend tracking
cmd_dict = validated_command.as_dict()
cmd_dict["_meta"] = {
"tokens_per_sec": round(tokens_per_sec, 2),
"latency_sec": round(elapsed_time, 2)
}
return cmd_dict
except (json.JSONDecodeError, CommandValidationError) as error:
error_msg = f"ERROR: Invalid command output: {error}\nRaw Output:\n{output_text}\nRepaired:\n{repaired}"
print(f"[Commander] {error_msg}")
if self.evaluator:
self.evaluator.log_llm_generation(tokens_generated, elapsed_time, False)
return {"error": error_msg}