#!/usr/bin/env python3 """ 虫群v8 — 推理服务适配器 让SwarmNode对外提供真实的参数化记忆模型推理服务 聚合协议的临时服务器通过此适配器调度各节点的模型 """ import time import hashlib from typing import Dict, List, Optional from core.parametric_memory import ParametricMemoryModel from core.aggregation_protocol.types import NodeInfo class InferenceService: """ 推理服务 — 包装参数化记忆模型,提供标准化推理接口 每个SwarmNode运行一个InferenceService 对外暴露:infer(query) → {response, confidence, latency_ms} """ def __init__(self, memory_model: ParametricMemoryModel, node_id: str): self.memory = memory_model self.node_id = node_id self._request_count = 0 self._total_latency_ms = 0.0 def infer(self, query: str, max_tokens: int = 64, temperature: float = 0.3) -> Dict: """ 推理入口 — 从参数化记忆模型生成回答 这就是"模型即数据库"的核心体现: 不查数据库,直接从模型参数生成个性化回答 """ start = time.time() self._request_count += 1 try: result = self.memory.recall( query, max_tokens=max_tokens, temperature=temperature ) latency = (time.time() - start) * 1000 self._total_latency_ms += latency return { "node_id": self.node_id, "response": result.get("response", ""), "confidence": result.get("confidence", 0.0), "source": result.get("source", "parametric"), "latency_ms": latency, "success": True, "error": "", } except Exception as e: latency = (time.time() - start) * 1000 return { "node_id": self.node_id, "response": "", "confidence": 0.0, "source": "error", "latency_ms": latency, "success": False, "error": str(e), } def store_and_learn(self, user_input: str, ai_response: str, memory_type: str = "chat", importance: float = 0.5) -> Dict: """ 存储记忆并触发学习 即时写入模式:存完即学,1-3步更新 """ start = time.time() mid = self.memory.store( user_input, ai_response, memory_type=memory_type, importance=importance, ) latency = (time.time() - start) * 1000 return { "memory_id": mid, "write_mode": self.memory.write_mode, "latency_ms": latency, "success": True, } def get_stats(self) -> Dict: """服务统计""" return { "node_id": self.node_id, "request_count": self._request_count, "avg_latency_ms": self._total_latency_ms / max(self._request_count, 1), "memory_status": self.memory.get_status(), } class DistributedInferenceBridge: """ 分布式推理桥 — 聚合协议与推理服务之间的桥梁 核心流程: 1. 聚合协议组建临时服务器 2. 桥接器将任务分发给各节点的InferenceService 3. 收集结果,按策略聚合 4. 返回最终结果 这让聚合协议能调度真实的参数化记忆模型 """ def __init__(self): # node_id -> InferenceService self._services: Dict[str, InferenceService] = {} def register_service(self, node_id: str, service: InferenceService): """注册节点的推理服务""" self._services[node_id] = service def unregister_service(self, node_id: str): """注销节点的推理服务""" self._services.pop(node_id, None) def dispatch_inference(self, query: str, node_ids: List[str], max_tokens: int = 64, temperature: float = 0.3) -> Dict: """ 分发推理到多个节点 1. 向所有指定节点发送推理请求 2. 收集结果 3. 聚合(置信度优先 + 内容互补) """ results = [] for nid in node_ids: service = self._services.get(nid) if not service: results.append({ "node_id": nid, "response": "", "confidence": 0.0, "success": False, "error": "服务不可用", }) continue result = service.infer(query, max_tokens=max_tokens, temperature=temperature) results.append(result) # 聚合 return self._aggregate(query, results) def _aggregate(self, query: str, results: List[Dict]) -> Dict: """ 聚合多节点推理结果 策略: 1. 过滤失败结果 2. 按置信度排序 3. 取最高置信度的作为主回答 4. 其他结果作为备选 5. 如果多个结果内容相似,提升置信度 """ successful = [r for r in results if r.get("success") and r.get("response")] if not successful: return { "response": "", "confidence": 0.0, "source": "distributed", "node_count": len(results), "successful_count": 0, "method": "none", "alternatives": [], } # 按置信度排序 successful.sort(key=lambda r: r.get("confidence", 0), reverse=True) primary = successful[0] # 内容相似度提升 confidence_boost = 0.0 alternatives = [] for r in successful[1:]: alternatives.append({ "node_id": r["node_id"], "response": r["response"][:100], "confidence": r.get("confidence", 0), }) # 简单相似度:关键词重叠 if self._is_similar(primary["response"], r["response"]): confidence_boost += 0.1 final_confidence = min(primary.get("confidence", 0) + confidence_boost, 1.0) return { "response": primary["response"], "confidence": final_confidence, "source": "distributed", "primary_node": primary["node_id"], "node_count": len(results), "successful_count": len(successful), "method": "confidence_with_consensus", "alternatives": alternatives[:3], } def _is_similar(self, text_a: str, text_b: str) -> bool: """简单相似度检测""" if not text_a or not text_b: return False # 取关键词集合 words_a = set(text_a[:50]) words_b = set(text_b[:50]) overlap = len(words_a & words_b) return overlap > len(words_a) * 0.3 def get_all_stats(self) -> Dict: """所有服务统计""" return { "total_services": len(self._services), "services": { nid: svc.get_stats() for nid, svc in self._services.items() }, }