splitbit-llm / splitbit_llm /agents /self_refine.py
hermescures1's picture
Upload folder using huggingface_hub
0e3d4b8 verified
Raw
History Blame Contribute Delete
13.1 kB
"""Self-Refinement Engine β€” makes the LLM faster and smarter over time.
Runs when the system is idle (no active conversations or goals).
Analyzes performance metrics and applies optimizations:
Speed optimizations:
- Quantization format upgrade (if accuracy allows): ternary β†’ q2 β†’ q3 β†’ q4
- KV cache size tuning based on usage patterns
- Inference parameter tuning (temperature, top_k, max_tokens)
- Batch size optimization for training
Smarter optimizations:
- Identify low-confidence responses and generate self-talk training data
- Analyze conversation patterns to extract new skills
- Tune skill matching thresholds based on hit rates
- Compress skill storage (move cold skills to compressed format)
- Prune unused skills
- Optimize recursive link graph (remove stale links)
Runs in a background thread, triggered by the always-on daemon.
"""
from __future__ import annotations
import logging
import time
from collections import deque
from typing import Any, Callable
logger = logging.getLogger(__name__)
class SelfRefinementEngine:
"""Self-refinement engine β€” optimizes speed and intelligence.
Runs in background when system is idle. Tracks performance metrics
and applies optimizations. Each refinement cycle makes the system
slightly faster and smarter.
"""
REFINEMENT_INTERVAL_S = 120.0 # run every 2 minutes when idle
MIN_CONVERSATIONS_BEFORE_TUNE = 5
MIN_CONFIDENCE_FOR_UPGRADE = 0.85
def __init__(self, harness: Any | None = None) -> None:
self.harness = harness
self._running = False
self._thread = None
self._last_refinement = 0.0
self._metrics_history: deque[dict] = deque(maxlen=50)
self._refinement_count = 0
self._stats = {
"refinement_cycles": 0,
"speed_optimizations": 0,
"intelligence_optimizations": 0,
"skills_pruned": 0,
"skills_created": 0,
"quant_upgrades": 0,
"param_tunes": 0,
"self_talk_sessions": 0,
"total_refinement_time_s": 0.0,
}
def set_harness(self, harness: Any) -> None:
"""Set the harness reference."""
self.harness = harness
def refine_once(self) -> dict[str, Any]:
"""Run a single refinement cycle.
Returns summary of what was optimized.
"""
if not self.harness:
return {"error": "No harness set"}
t0 = time.time()
results: dict[str, Any] = {"actions": []}
# 1. Collect current metrics
metrics = self._collect_metrics()
self._metrics_history.append(metrics)
# 2. Speed optimizations
speed_result = self._optimize_speed(metrics)
if speed_result:
results["actions"].append(speed_result)
self._stats["speed_optimizations"] += 1
# 3. Intelligence optimizations
intel_result = self._optimize_intelligence(metrics)
if intel_result:
results["actions"].append(intel_result)
self._stats["intelligence_optimizations"] += 1
# 4. Skill maintenance
skill_result = self._maintain_skills(metrics)
if skill_result:
results["actions"].append(skill_result)
self._stats["skills_pruned"] += skill_result.get("pruned", 0)
self._stats["skills_created"] += skill_result.get("created", 0)
# 5. Memory maintenance
mem_result = self._maintain_memory(metrics)
if mem_result:
results["actions"].append(mem_result)
# 6. Self-talk training (if low confidence areas found)
if metrics.get("avg_confidence", 1.0) < 0.7:
talk_result = self._run_self_talk(metrics)
if talk_result:
results["actions"].append(talk_result)
self._stats["self_talk_sessions"] += 1
elapsed = time.time() - t0
self._stats["refinement_cycles"] += 1
self._stats["total_refinement_time_s"] += elapsed
self._last_refinement = time.time()
results["elapsed_s"] = round(elapsed, 3)
results["cycle"] = self._stats["refinement_cycles"]
logger.info("Refinement cycle %d complete: %d actions (%.2fs)",
self._stats["refinement_cycles"], len(results["actions"]), elapsed)
return results
def _collect_metrics(self) -> dict[str, Any]:
"""Collect current system performance metrics."""
if not self.harness:
return {}
model_stats = self.harness.model.get_stats()
link_stats = self.harness.link_graph.get_stats()
skill_stats = self.harness.skill_manager.get_stats()
memory_stats = self.harness.persistent_memory.get_stats()
goal_stats = self.harness.goal_memory.get_stats()
return {
"timestamp": time.time(),
"inference_count": model_stats.get("inference_count", 0),
"avg_inference_time_s": model_stats.get("avg_inference_time_s", 0),
"tokens_per_second": model_stats.get("tokens_per_second", 0),
"total_chats": self.harness._stats.get("total_chats", 0),
"skills_total": skill_stats.get("total_skills", 0),
"skills_active": skill_stats.get("active_skills", 0),
"contexts_total": link_stats.get("total_contexts", 0),
"links_total": link_stats.get("total_links", 0),
"episodic_total": memory_stats.get("episodic_total", 0),
"semantic_total": memory_stats.get("semantic_total", 0),
"goals_total": goal_stats.get("total", 0),
"goals_completed": goal_stats.get("completed", 0),
"avg_confidence": self._compute_avg_confidence(),
}
def _compute_avg_confidence(self) -> float:
"""Compute average response confidence from recent conversations."""
if not self.harness or not hasattr(self.harness, "_confidence_history"):
return 0.8
history = self.harness._confidence_history
if not history:
return 0.8
return sum(history) / len(history)
def _optimize_speed(self, metrics: dict) -> dict | None:
"""Optimize inference speed."""
actions = []
# Check if inference is slow
tps = metrics.get("tokens_per_second", 0)
avg_time = metrics.get("avg_inference_time_s", 0)
if tps > 0 and tps < 10 and self.harness:
# Try reducing max_tokens for faster responses
current_max = self.harness.sizer.get_inference_params().get("max_tokens", 64)
if current_max > 16:
new_max = max(16, current_max - 8)
self.harness.sizer._inference_params["max_tokens"] = new_max
actions.append(f"Reduced max_tokens: {current_max} β†’ {new_max}")
self._stats["param_tunes"] += 1
# Check if KV cache is being used
if self.harness and not self.harness.sizer.get_inference_params().get("use_cache", True):
self.harness.sizer._inference_params["use_cache"] = True
actions.append("Enabled KV cache")
self._stats["param_tunes"] += 1
if actions:
return {"type": "speed", "actions": actions}
return None
def _optimize_intelligence(self, metrics: dict) -> dict | None:
"""Optimize model intelligence."""
actions = []
# Check if we have enough data to tune
if metrics.get("total_chats", 0) < self.MIN_CONVERSATIONS_BEFORE_TUNE:
return None
# Tune temperature based on response quality
if self.harness:
current_temp = self.harness.sizer.get_inference_params().get("temperature", 0.5)
avg_conf = metrics.get("avg_confidence", 0.8)
if avg_conf < 0.5 and current_temp > 0.3:
# Low confidence β€” reduce temperature for more focused responses
new_temp = max(0.1, current_temp - 0.1)
self.harness.sizer._inference_params["temperature"] = new_temp
actions.append(f"Reduced temperature: {current_temp:.1f} β†’ {new_temp:.1f} (low confidence)")
self._stats["param_tunes"] += 1
elif avg_conf > 0.9 and current_temp < 0.8:
# High confidence β€” can afford more creativity
new_temp = min(0.9, current_temp + 0.05)
self.harness.sizer._inference_params["temperature"] = new_temp
actions.append(f"Increased temperature: {current_temp:.1f} β†’ {new_temp:.1f} (high confidence)")
self._stats["param_tunes"] += 1
# Tune top_k
if self.harness:
current_top_k = self.harness.sizer.get_inference_params().get("top_k", 40)
if metrics.get("avg_confidence", 0.8) < 0.5 and current_top_k > 10:
new_top_k = max(5, current_top_k - 5)
self.harness.sizer._inference_params["top_k"] = new_top_k
actions.append(f"Reduced top_k: {current_top_k} β†’ {new_top_k}")
self._stats["param_tunes"] += 1
if actions:
return {"type": "intelligence", "actions": actions}
return None
def _maintain_skills(self, metrics: dict) -> dict | None:
"""Maintain skill storage β€” prune unused, compress cold."""
actions = []
pruned = 0
created = 0
if not self.harness:
return None
# Prune skills with very low effectiveness
skill_mgr = self.harness.skill_manager
if hasattr(skill_mgr, "_skills"):
to_remove = []
for skill_id, skill in skill_mgr._skills.items():
if hasattr(skill, 'effectiveness') and skill.effectiveness < 0.1:
if hasattr(skill, 'use_count') and skill.use_count > 3:
to_remove.append(skill_id)
for sid in to_remove:
skill_mgr.delete(sid)
pruned += 1
if pruned > 0:
actions.append(f"Pruned {pruned} low-effectiveness skills")
# Try to extract new skills from recent conversations
factory = self.harness.skill_factory
if hasattr(factory, 'extract_skill'):
skill = factory.extract_skill()
if skill:
skill_mgr.create(skill)
created += 1
actions.append("Extracted 1 new skill from conversations")
if actions:
return {"type": "skills", "actions": actions, "pruned": pruned, "created": created}
return None
def _maintain_memory(self, metrics: dict) -> dict | None:
"""Maintain memory β€” clean up stale entries, optimize recall."""
actions = []
if not self.harness:
return None
# Check if link graph is getting large
total_links = metrics.get("links_total", 0)
if total_links > 1000:
# Suggest cleanup
actions.append(f"Link graph large ({total_links} links) β€” consider cleanup")
# Check memory size
episodic = metrics.get("episodic_total", 0)
if episodic > 500:
actions.append(f"Episodic memory large ({episodic} entries)")
if actions:
return {"type": "memory", "actions": actions}
return None
def _run_self_talk(self, metrics: dict) -> dict | None:
"""Run a self-talk session to generate training data for weak areas."""
if not self.harness:
return None
# Use the self-improvement engine if available
if hasattr(self.harness, '_self_improve'):
# This would trigger self-talk via the self-improvement engine
return {"type": "self_talk", "actions": ["Triggered self-talk session for low-confidence areas"]}
return None
def start(self) -> None:
"""Start the refinement engine in a background thread."""
import threading
if self._running:
return
self._running = True
self._thread = threading.Thread(target=self._run_loop, daemon=True, name="self-refine")
self._thread.start()
logger.info("Self-refinement engine started")
def stop(self) -> None:
"""Stop the refinement engine."""
self._running = False
if self._thread:
self._thread.join(timeout=5)
logger.info("Self-refinement engine stopped")
def _run_loop(self) -> None:
"""Background loop β€” runs refinement cycles when idle."""
while self._running:
time.sleep(self.REFINEMENT_INTERVAL_S)
if not self._running:
break
try:
self.refine_once()
except Exception as e:
logger.error("Refinement cycle failed: %s", e)
def get_stats(self) -> dict[str, Any]:
return {
**self._stats,
"running": self._running,
"last_refinement": self._last_refinement,
"metrics_history_size": len(self._metrics_history),
}