diff --git a/.env.example b/.env.example index 834887c..0346921 100644 --- a/.env.example +++ b/.env.example @@ -6,5 +6,15 @@ OPENAI_API_KEY= # platform.openai.com GOOGLE_API_KEY= # aistudio.google.com/api-keys NVIDIA_API_KEY= # build.nvidia.com +# Third-party providers (optional) +SILICONFLOW_API_KEY= # siliconflow.cn +OPENROUTER_API_KEY= # openrouter.ai +ZHIPU_API_KEY= # open.bigmodel.cn +CUSTOM_API_KEY= # Your custom OpenAI-compatible endpoint +CUSTOM_BASE_URL= # Your custom API base URL (optional) + +# Local models (optional) +OLLAMA_BASE_URL= # http://localhost:11434 (default) + # Web search (optional) TAVILY_API_KEY= # app.tavily.com diff --git a/.gitignore b/.gitignore index ea24a53..256d8e9 100644 --- a/.gitignore +++ b/.gitignore @@ -10,6 +10,7 @@ build/ *.egg *.pytest_cache/ .coverage +.ipynb_checkpoints/ # Environment .env @@ -41,3 +42,4 @@ media/ *AGENTS.md *meals/ botpy.log +large_tool_results/ \ No newline at end of file diff --git a/EvoScientist/config/onboard.py b/EvoScientist/config/onboard.py index 7c6978e..6a74e50 100644 --- a/EvoScientist/config/onboard.py +++ b/EvoScientist/config/onboard.py @@ -35,26 +35,30 @@ console = Console() # Wizard Style # ============================================================================= -WIZARD_STYLE = Style.from_dict({ - "qmark": "fg:#00bcd4 bold", # Cyan question mark - "question": "bold", # Bold question text - "answer": "fg:#4caf50 bold", # Green selected answer - "pointer": "fg:#4caf50", # Green pointer (») - "highlighted": "noreverse bold", # No background, bold text - "selected": "fg:#4caf50 bold", # Green ● indicator - "separator": "fg:#6c6c6c", # Dim separator - "disabled": "fg:#858585", # Dim disabled indicator (-) - "instruction": "fg:#858585", # Dim instructions - "text": "fg:#858585", # Dim gray ○ and unselected text -}) +WIZARD_STYLE = Style.from_dict( + { + "qmark": "fg:#00bcd4 bold", # Cyan question mark + "question": "bold", # Bold question text + "answer": "fg:#4caf50 bold", # Green selected answer + "pointer": "fg:#4caf50", # Green pointer (») + "highlighted": "noreverse bold", # No background, bold text + "selected": "fg:#4caf50 bold", # Green ● indicator + "separator": "fg:#6c6c6c", # Dim separator + "disabled": "fg:#858585", # Dim disabled indicator (-) + "instruction": "fg:#858585", # Dim instructions + "text": "fg:#858585", # Dim gray ○ and unselected text + } +) -CONFIRM_STYLE = Style.from_dict({ - "qmark": "fg:#e69500 bold", # Orange warning mark (!) - "question": "bold", - "answer": "fg:#4caf50 bold", - "instruction": "fg:#858585", - "text": "", -}) +CONFIRM_STYLE = Style.from_dict( + { + "qmark": "fg:#e69500 bold", # Orange warning mark (!) + "question": "bold", + "answer": "fg:#4caf50 bold", + "instruction": "fg:#858585", + "text": "", + } +) QMARK = "❯" @@ -85,18 +89,35 @@ def _checkbox_ask(choices, message: str, **kwargs): InquirerControl._get_choice_tokens = _patched try: return questionary.checkbox( - message, choices=choices, style=WIZARD_STYLE, qmark=QMARK, **kwargs, + message, + choices=choices, + style=WIZARD_STYLE, + qmark=QMARK, + **kwargs, ).ask() finally: InquirerControl._get_choice_tokens = original -STEPS = ["UI", "Provider", "API Key", "Model", "Tavily Key", "Workspace", "Thinking", "Skills", "MCP Servers", "Channels"] + +STEPS = [ + "UI", + "Provider", + "API Key", + "Model", + "Tavily Key", + "Workspace", + "Thinking", + "Skills", + "MCP Servers", + "Channels", +] # ============================================================================= # Validators # ============================================================================= + class IntegerValidator(Validator): """Validates that input is a positive integer.""" @@ -130,15 +151,14 @@ class ChoiceValidator(Validator): if not text and self.allow_empty: return if text not in [c.lower() for c in self.choices]: - raise ValidationError( - message=f"Must be one of: {', '.join(self.choices)}" - ) + raise ValidationError(message=f"Must be one of: {', '.join(self.choices)}") # ============================================================================= # API Key Validation # ============================================================================= + def validate_anthropic_key(api_key: str) -> tuple[bool, str]: """Validate an Anthropic API key by making a test request. @@ -153,6 +173,7 @@ def validate_anthropic_key(api_key: str) -> tuple[bool, str]: try: import anthropic + client = anthropic.Anthropic(api_key=api_key) # Make a minimal request to validate the key client.models.list() @@ -177,6 +198,7 @@ def validate_openai_key(api_key: str) -> tuple[bool, str]: try: import openai + client = openai.OpenAI(api_key=api_key) # Make a minimal request to validate the key client.models.list() @@ -201,12 +223,18 @@ def validate_nvidia_key(api_key: str) -> tuple[bool, str]: try: from langchain_nvidia_ai_endpoints import ChatNVIDIA + llm = ChatNVIDIA(api_key=api_key, model="meta/llama-3.1-8b-instruct") llm.available_models return True, "Valid" except Exception as e: error_str = str(e).lower() - if "401" in error_str or "unauthorized" in error_str or "invalid" in error_str or "authentication" in error_str: + if ( + "401" in error_str + or "unauthorized" in error_str + or "invalid" in error_str + or "authentication" in error_str + ): return False, "Invalid API key" return False, f"Error: {e}" @@ -225,6 +253,7 @@ def validate_google_key(api_key: str) -> tuple[bool, str]: try: from google import genai + client = genai.Client(api_key=api_key) # Make a minimal request to validate the key pager = client.models.list(config={"page_size": 1}) @@ -235,7 +264,14 @@ def validate_google_key(api_key: str) -> tuple[bool, str]: return True, "Valid" except Exception as e: error_str = str(e).lower() - if "400" in error_str or "401" in error_str or "403" in error_str or "unauthorized" in error_str or "invalid" in error_str or "api key" in error_str: + if ( + "400" in error_str + or "401" in error_str + or "403" in error_str + or "unauthorized" in error_str + or "invalid" in error_str + or "api key" in error_str + ): return False, "Invalid API key" return False, f"Error: {e}" @@ -251,12 +287,20 @@ def validate_siliconflow_key(api_key: str) -> tuple[bool, str]: try: import openai - client = openai.OpenAI(api_key=api_key, base_url="https://api.siliconflow.cn/v1") + + client = openai.OpenAI( + api_key=api_key, base_url="https://api.siliconflow.cn/v1" + ) client.models.list() return True, "Valid" except Exception as e: error_str = str(e).lower() - if "401" in error_str or "unauthorized" in error_str or "invalid" in error_str or "authentication" in error_str: + if ( + "401" in error_str + or "unauthorized" in error_str + or "invalid" in error_str + or "authentication" in error_str + ): return False, "Invalid API key" return False, f"Error: {e}" @@ -272,12 +316,47 @@ def validate_openrouter_key(api_key: str) -> tuple[bool, str]: try: import openai + client = openai.OpenAI(api_key=api_key, base_url="https://openrouter.ai/api/v1") client.models.list() return True, "Valid" except Exception as e: error_str = str(e).lower() - if "401" in error_str or "unauthorized" in error_str or "invalid" in error_str or "authentication" in error_str: + if ( + "401" in error_str + or "unauthorized" in error_str + or "invalid" in error_str + or "authentication" in error_str + ): + return False, "Invalid API key" + return False, f"Error: {e}" + + +def validate_zhipu_key(api_key: str) -> tuple[bool, str]: + """Validate a ZhipuAI API key by making a test request. + + Returns: + Tuple of (is_valid, message). + """ + if not api_key: + return True, "Skipped (no key provided)" + + try: + import openai + + client = openai.OpenAI( + api_key=api_key, base_url="https://open.bigmodel.cn/api/paas/v4" + ) + client.models.list() + return True, "Valid" + except Exception as e: + error_str = str(e).lower() + if ( + "401" in error_str + or "unauthorized" in error_str + or "invalid" in error_str + or "authentication" in error_str + ): return False, "Invalid API key" return False, f"Error: {e}" @@ -296,6 +375,7 @@ def validate_tavily_key(api_key: str) -> tuple[bool, str]: try: from tavily import TavilyClient + client = TavilyClient(api_key=api_key) # Make a minimal search to validate client.search("test", max_results=1) @@ -322,6 +402,7 @@ def validate_ollama_connection(base_url: str) -> tuple[bool, str, list[str]]: try: import httpx + resp = httpx.get(f"{base_url.rstrip('/')}/api/tags", timeout=5) if resp.status_code == 200: data = resp.json() @@ -340,17 +421,20 @@ def validate_ollama_connection(base_url: str) -> tuple[bool, str, list[str]]: # Display Helpers # ============================================================================= + def _print_header() -> None: """Print the wizard header.""" console.print() - console.print(Panel.fit( - Text.from_markup( - "[bold cyan]EvoScientist Setup Wizard[/bold cyan]\n\n" - "This wizard will help you configure EvoScientist.\n" - "Press Ctrl+C at any time to cancel." - ), - border_style="cyan", - )) + console.print( + Panel.fit( + Text.from_markup( + "[bold cyan]EvoScientist Setup Wizard[/bold cyan]\n\n" + "This wizard will help you configure EvoScientist.\n" + "Press Ctrl+C at any time to cancel." + ), + border_style="cyan", + ) + ) console.print() @@ -380,6 +464,7 @@ def _print_step_skipped(step_name: str, reason: str = "kept current") -> None: # Step Functions # ============================================================================= + def _step_ui_backend(config: EvoScientistConfig) -> str: """Step 0: Select UI backend (Rich CLI or Textual TUI). @@ -423,14 +508,41 @@ def _step_provider(config: EvoScientistConfig) -> str: Choice(title="OpenAI (GPT models)", value="openai"), Choice(title="Google GenAI (Gemini models)", value="google-genai"), Choice(title="NVIDIA (third party — limited free requests)", value="nvidia"), - Choice(title="SiliconFlow (third party — GLM, Kimi, MiniMax, etc.)", value="siliconflow"), - Choice(title="OpenRouter (third party — Grok, Gemini, Qwen, etc.)", value="openrouter"), + Choice( + title="SiliconFlow (third party — GLM, Kimi, MiniMax, etc.)", + value="siliconflow", + ), + Choice( + title="OpenRouter (third party — Grok, Gemini, Qwen, etc.)", + value="openrouter", + ), + Choice(title="ZhipuAI (智谱 — GLM models)", value="zhipu"), + Choice( + title="ZhipuAI CodePlan (智谱代码计划 — GLM models for coding)", + value="zhipu-code", + ), Choice(title="Ollama (local models)", value="ollama"), Choice(title="Other (OpenAI-compatible)", value="custom"), ] # Set default based on current config - default = config.provider if config.provider in ["anthropic", "openai", "google-genai", "nvidia", "siliconflow", "openrouter", "ollama", "custom"] else "anthropic" + default = ( + config.provider + if config.provider + in [ + "anthropic", + "openai", + "google-genai", + "nvidia", + "siliconflow", + "openrouter", + "zhipu", + "zhipu-code", + "ollama", + "custom", + ] + else "anthropic" + ) provider = questionary.select( "Select your LLM provider:", @@ -450,15 +562,56 @@ def _step_provider(config: EvoScientistConfig) -> str: def _provider_key_info(config: EvoScientistConfig, provider: str): """Return (display_name, current_value, validate_fn) for a provider.""" mapping = { - "anthropic": ("Anthropic", config.anthropic_api_key or os.environ.get("ANTHROPIC_API_KEY", ""), validate_anthropic_key), - "nvidia": ("NVIDIA", config.nvidia_api_key or os.environ.get("NVIDIA_API_KEY", ""), validate_nvidia_key), - "google-genai": ("Google", config.google_api_key or os.environ.get("GOOGLE_API_KEY", ""), validate_google_key), - "siliconflow": ("SiliconFlow", config.siliconflow_api_key or os.environ.get("SILICONFLOW_API_KEY", ""), validate_siliconflow_key), - "openrouter": ("OpenRouter", config.openrouter_api_key or os.environ.get("OPENROUTER_API_KEY", ""), validate_openrouter_key), - "custom": ("Custom", config.custom_api_key or os.environ.get("CUSTOM_API_KEY", ""), None), - "ollama": ("Ollama", "__no_key__", None), + "anthropic": ( + "Anthropic", + config.anthropic_api_key or os.environ.get("ANTHROPIC_API_KEY", ""), + validate_anthropic_key, + ), + "nvidia": ( + "NVIDIA", + config.nvidia_api_key or os.environ.get("NVIDIA_API_KEY", ""), + validate_nvidia_key, + ), + "google-genai": ( + "Google", + config.google_api_key or os.environ.get("GOOGLE_API_KEY", ""), + validate_google_key, + ), + "siliconflow": ( + "SiliconFlow", + config.siliconflow_api_key or os.environ.get("SILICONFLOW_API_KEY", ""), + validate_siliconflow_key, + ), + "openrouter": ( + "OpenRouter", + config.openrouter_api_key or os.environ.get("OPENROUTER_API_KEY", ""), + validate_openrouter_key, + ), + "zhipu": ( + "ZhipuAI", + config.zhipu_api_key or os.environ.get("ZHIPU_API_KEY", ""), + validate_zhipu_key, + ), + "zhipu-code": ( + "ZhipuAI CodePlan", + config.zhipu_api_key or os.environ.get("ZHIPU_API_KEY", ""), + validate_zhipu_key, + ), + "custom": ( + "Custom", + config.custom_api_key or os.environ.get("CUSTOM_API_KEY", ""), + None, + ), + "ollama": ("Ollama", "__no_key__", None), } - return mapping.get(provider, ("OpenAI", config.openai_api_key or os.environ.get("OPENAI_API_KEY", ""), validate_openai_key)) + return mapping.get( + provider, + ( + "OpenAI", + config.openai_api_key or os.environ.get("OPENAI_API_KEY", ""), + validate_openai_key, + ), + ) def _prompt_and_validate_api_key( @@ -541,7 +694,10 @@ def _step_provider_api_key( prompt_text = f"Enter {key_name} API key ({hint}, Enter to keep):" return _prompt_and_validate_api_key( - prompt_text, current, validate_fn, skip_validation, + prompt_text, + current, + validate_fn, + skip_validation, ) @@ -563,7 +719,9 @@ def _step_base_url(config: EvoScientistConfig) -> str: default=default, style=WIZARD_STYLE, qmark=QMARK, - placeholder=FormattedText([("fg:#858585", " e.g. https://api.example.com/v1")]) if not default else None, + placeholder=FormattedText([("fg:#858585", " e.g. https://api.example.com/v1")]) + if not default + else None, ).ask() if url is None: raise KeyboardInterrupt() @@ -626,8 +784,7 @@ def _step_model( if ollama_detected_models: _CUSTOM_SENTINEL = "__custom__" choices = [ - Choice(title=name, value=name) - for name in ollama_detected_models + Choice(title=name, value=name) for name in ollama_detected_models ] choices.append(Choice(title="Type a model name...", value=_CUSTOM_SENTINEL)) @@ -651,7 +808,9 @@ def _step_model( # No detected models (server down or empty) — direct text input if not ollama_detected_models: - console.print(" [dim]No models detected — type the model name you plan to pull.[/dim]") + console.print( + " [dim]No models detected — type the model name you plan to pull.[/dim]" + ) model = questionary.text( "Model name:", style=WIZARD_STYLE, @@ -745,7 +904,10 @@ def _step_tavily_key( prompt_text = f"Tavily API key for web search ({hint}, Enter to keep):" return _prompt_and_validate_api_key( - prompt_text, current, validate_tavily_key, skip_validation, + prompt_text, + current, + validate_tavily_key, + skip_validation, placeholder=FormattedText([("fg:#858585", " (recommended for web search)")]), ) @@ -875,7 +1037,9 @@ def _check_npx() -> bool: try: result = subprocess.run( ["npx", "--version"], - capture_output=True, text=True, timeout=10, + capture_output=True, + text=True, + timeout=10, ) return result.returncode == 0 except (FileNotFoundError, subprocess.TimeoutExpired): @@ -897,7 +1061,9 @@ def _detect_node_install_method() -> tuple[str, str]: try: result = subprocess.run( ["brew", "--version"], - capture_output=True, text=True, timeout=5, + capture_output=True, + text=True, + timeout=5, ) if result.returncode == 0: return "brew", "brew install node" @@ -919,7 +1085,8 @@ def _install_node(method: str, command: str) -> bool: try: proc = subprocess.run( command.split(), - capture_output=True, text=True, + capture_output=True, + text=True, timeout=120, ) return proc.returncode == 0 @@ -965,7 +1132,9 @@ def _ensure_npx(reason: str) -> bool: console.print(" [green]✓ npx now available[/green]") return True else: - console.print(" [yellow]✗ npx still not found after install[/yellow]") + console.print( + " [yellow]✗ npx still not found after install[/yellow]" + ) else: console.print(" [red]✗ Installation failed[/red]") else: @@ -1016,11 +1185,12 @@ def _step_skills() -> list[str]: choices.append(Choice(title=skill["label"], value=skill["source"])) all_installed = all( - _hint_name(skill["source"]) in installed_names - for skill in _RECOMMENDED_SKILLS + _hint_name(skill["source"]) in installed_names for skill in _RECOMMENDED_SKILLS ) if all_installed: - console.print(" [green]✓ All recommended skills are already installed.[/green]") + console.print( + " [green]✓ All recommended skills are already installed.[/green]" + ) return [] selected = _checkbox_ask(choices, "Install or Sync predefined skills:") @@ -1035,7 +1205,9 @@ def _step_skills() -> list[str]: if has_npx: _print_step_skipped("Skills", "none selected — good choice!") console.print(" [green]✓ npx found — skill discovery available[/green]") - console.print(" [yellow bold]* Less is more[/yellow bold] [dim](EvoScientist can discover and install skills on its own)[/dim]") + console.print( + " [yellow bold]* Less is more[/yellow bold] [dim](EvoScientist can discover and install skills on its own)[/dim]" + ) else: _print_step_skipped("Skills", "none selected") @@ -1052,7 +1224,9 @@ def _step_skills() -> list[str]: _print_step_result("Skill", label) installed.append(source) else: - _print_step_result("Skill", f"{label} — {result.get('error', 'failed')}", success=False) + _print_step_result( + "Skill", f"{label} — {result.get('error', 'failed')}", success=False + ) except Exception as e: _print_step_result("Skill", f"{label} — {e}", success=False) @@ -1117,7 +1291,9 @@ def _install_pip_package(package: str) -> bool: try: result = subprocess.run( [sys.executable, "-m", "pip", "install", "-q", package], - capture_output=True, text=True, timeout=120, + capture_output=True, + text=True, + timeout=120, ) return result.returncode == 0 except (FileNotFoundError, subprocess.TimeoutExpired): @@ -1156,9 +1332,13 @@ def _step_mcp_servers() -> list[str]: else: choices.append(Choice(title=srv["label"], value=srv["name"])) - all_installed = all(srv["name"] in existing_config for srv in _RECOMMENDED_MCP_SERVERS) + all_installed = all( + srv["name"] in existing_config for srv in _RECOMMENDED_MCP_SERVERS + ) if all_installed: - console.print(" [green]✓ All recommended MCP servers are already configured.[/green]") + console.print( + " [green]✓ All recommended MCP servers are already configured.[/green]" + ) return [] selected = _checkbox_ask(choices, "Install recommended MCP servers:") @@ -1168,7 +1348,9 @@ def _step_mcp_servers() -> list[str]: if not selected: _print_step_skipped("MCP Servers", "none selected") - console.print(" [dim]Add later with: EvoSci mcp add [--env-ref KEY] -- [args][/dim]") + console.print( + " [dim]Add later with: EvoSci mcp add [--env-ref KEY] -- [args][/dim]" + ) return [] # Check if any selected servers require npx @@ -1186,7 +1368,9 @@ def _step_mcp_servers() -> list[str]: } selected = [s for s in selected if s not in npx_servers] if npx_servers: - console.print(f" [yellow]\u26a0 Skipping {', '.join(sorted(npx_servers))} (npx not available)[/yellow]") + console.print( + f" [yellow]\u26a0 Skipping {', '.join(sorted(npx_servers))} (npx not available)[/yellow]" + ) if not selected: return [] @@ -1205,14 +1389,18 @@ def _step_mcp_servers() -> list[str]: console.print(f" [yellow]⚠ Requires {env_key}[/yellow]") console.print(f" [dim]{hint}[/dim]") if not os.environ.get(env_key): - console.print(f" [dim]Set it before running EvoScientist: export {env_key}=...[/dim]") + console.print( + f" [dim]Set it before running EvoScientist: export {env_key}=...[/dim]" + ) # Install pip package if needed pip_pkg = srv.get("pip_package") if pip_pkg: console.print(f" [dim]Installing {pip_pkg}...[/dim]") if not _install_pip_package(pip_pkg): - _print_step_result("MCP", f"{name} — pip install {pip_pkg} failed", success=False) + _print_step_result( + "MCP", f"{name} — pip install {pip_pkg} failed", success=False + ) continue # Add to MCP config @@ -1220,7 +1408,8 @@ def _step_mcp_servers() -> list[str]: add_mcp_server(name, "streamable_http", url=srv["url"]) else: add_mcp_server( - name, "stdio", + name, + "stdio", command=srv["command"], args=srv.get("args", []), env=srv.get("env"), @@ -1253,7 +1442,9 @@ def validate_imessage() -> tuple[bool, str]: try: result = subprocess.run( [cli_path, "--version"], - capture_output=True, text=True, timeout=5, + capture_output=True, + text=True, + timeout=5, ) version = result.stdout.strip() if result.returncode == 0 else None except Exception: @@ -1263,14 +1454,19 @@ def validate_imessage() -> tuple[bool, str]: try: result = subprocess.run( [cli_path, "rpc", "--help"], - capture_output=True, text=True, timeout=5, + capture_output=True, + text=True, + timeout=5, ) rpc_ok = result.returncode == 0 except Exception: rpc_ok = False if not rpc_ok: - return False, f"imsg found at {cli_path} but RPC not supported (update with: brew upgrade imsg)" + return ( + False, + f"imsg found at {cli_path} but RPC not supported (update with: brew upgrade imsg)", + ) version_str = f" ({version})" if version else "" return True, f"imsg{version_str} at {cli_path}" @@ -1285,7 +1481,8 @@ def _install_imsg() -> bool: try: proc = subprocess.run( ["brew", "install", "steipete/tap/imsg"], - capture_output=True, text=True, + capture_output=True, + text=True, timeout=120, ) return proc.returncode == 0 @@ -1349,7 +1546,9 @@ def _setup_imessage() -> bool: else: return False else: - console.print(" [dim]Skipped. Install manually: brew install steipete/tap/imsg[/dim]") + console.print( + " [dim]Skipped. Install manually: brew install steipete/tap/imsg[/dim]" + ) return False else: # RPC not supported or other issue @@ -1378,22 +1577,97 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]: if t.strip() } # Legacy iMessage compat - if getattr(config, "imessage_enabled", False) and "imessage" not in _currently_enabled: + if ( + getattr(config, "imessage_enabled", False) + and "imessage" not in _currently_enabled + ): _currently_enabled.add("imessage") # Channel definitions: (value, display_name, required_fields, import_check, pip_extra) # import_check: module name to try importing; None = no check needed _CHANNELS = [ - ("telegram", "Telegram", [("telegram_bot_token", "Bot token (from @BotFather)")], "telegram", "telegram"), - ("discord", "Discord", [("discord_bot_token", "Bot token")], "discord", "discord"), - ("slack", "Slack", [("slack_bot_token", "Bot token (xoxb-...)"), ("slack_app_token", "App token for Socket Mode (xapp-...)")], "slack_sdk", "slack"), - ("feishu", "Feishu", [("feishu_app_id", "App ID"), ("feishu_app_secret", "App Secret")], "aiohttp", "feishu"), - ("dingtalk", "DingTalk", [("dingtalk_client_id", "Client ID (AppKey)"), ("dingtalk_client_secret", "Client Secret (AppSecret)")], "aiohttp", "dingtalk"), - ("wechat", "WeChat", [("wechat_wecom_corp_id", "WeCom Corp ID"), ("wechat_wecom_agent_id", "WeCom Agent ID"), ("wechat_wecom_secret", "WeCom Secret")], "aiohttp", "wechat"), - ("email", "Email", [("email_imap_host", "IMAP host"), ("email_imap_username", "IMAP username"), ("email_imap_password", "IMAP password"), ("email_smtp_host", "SMTP host"), ("email_smtp_username", "SMTP username"), ("email_smtp_password", "SMTP password"), ("email_from_address", "From address")], None, None), - ("qq", "QQ", [("qq_app_id", "App ID"), ("qq_app_secret", "App Secret")], "botpy", "qq"), - ("signal", "Signal", [("signal_phone_number", "Phone number (E.164)")], None, None), - ("imessage", "iMessage", [], None, None), # handled via _setup_imessage() + ( + "telegram", + "Telegram", + [("telegram_bot_token", "Bot token (from @BotFather)")], + "telegram", + "telegram", + ), + ( + "discord", + "Discord", + [("discord_bot_token", "Bot token")], + "discord", + "discord", + ), + ( + "slack", + "Slack", + [ + ("slack_bot_token", "Bot token (xoxb-...)"), + ("slack_app_token", "App token for Socket Mode (xapp-...)"), + ], + "slack_sdk", + "slack", + ), + ( + "feishu", + "Feishu", + [("feishu_app_id", "App ID"), ("feishu_app_secret", "App Secret")], + "aiohttp", + "feishu", + ), + ( + "dingtalk", + "DingTalk", + [ + ("dingtalk_client_id", "Client ID (AppKey)"), + ("dingtalk_client_secret", "Client Secret (AppSecret)"), + ], + "aiohttp", + "dingtalk", + ), + ( + "wechat", + "WeChat", + [ + ("wechat_wecom_corp_id", "WeCom Corp ID"), + ("wechat_wecom_agent_id", "WeCom Agent ID"), + ("wechat_wecom_secret", "WeCom Secret"), + ], + "aiohttp", + "wechat", + ), + ( + "email", + "Email", + [ + ("email_imap_host", "IMAP host"), + ("email_imap_username", "IMAP username"), + ("email_imap_password", "IMAP password"), + ("email_smtp_host", "SMTP host"), + ("email_smtp_username", "SMTP username"), + ("email_smtp_password", "SMTP password"), + ("email_from_address", "From address"), + ], + None, + None, + ), + ( + "qq", + "QQ", + [("qq_app_id", "App ID"), ("qq_app_secret", "App Secret")], + "botpy", + "qq", + ), + ( + "signal", + "Signal", + [("signal_phone_number", "Phone number (E.164)")], + None, + None, + ), + ("imessage", "iMessage", [], None, None), # handled via _setup_imessage() ] choices = [ @@ -1423,7 +1697,9 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]: return updates # Build a lookup for channel definitions - _ch_lookup = {v: (v, d, fields, imp, extra) for v, d, fields, imp, extra in _CHANNELS} + _ch_lookup = { + v: (v, d, fields, imp, extra) for v, d, fields, imp, extra in _CHANNELS + } enabled_channels: list[str] = [] @@ -1437,7 +1713,9 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]: __import__(import_check) except ImportError: console.print(" [yellow]✗ Required package not installed.[/yellow]") - console.print(f' [dim]Run:[/dim] pip install "evoscientist\\[{pip_extra}]"') + console.print( + f' [dim]Run:[/dim] pip install "evoscientist\\[{pip_extra}]"' + ) console.print(" [dim]Then re-run:[/dim] EvoSci onboard") continue @@ -1485,7 +1763,9 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]: # Feishu optional fields (verification_token & encrypt_key) if ch_name == "feishu": - console.print(" [dim]The following fields are optional (press Enter to skip):[/dim]") + console.print( + " [dim]The following fields are optional (press Enter to skip):[/dim]" + ) for field_name, prompt_label in [ ("feishu_verification_token", "Verification Token (optional)"), ("feishu_encrypt_key", "Encrypt Key (optional)"), @@ -1570,18 +1850,21 @@ def _probe_channel( async def _run() -> tuple[bool, str]: if ch_name == "telegram": from ..channels.telegram.probe import validate_telegram_token + return await validate_telegram_token( _val("telegram_bot_token"), _val("telegram_proxy") or None, ) elif ch_name == "discord": from ..channels.discord.probe import validate_discord_token + return await validate_discord_token( _val("discord_bot_token"), _val("discord_proxy") or None, ) elif ch_name == "slack": from ..channels.slack.probe import validate_slack_tokens + return await validate_slack_tokens( _val("slack_bot_token"), _val("slack_app_token") or None, @@ -1591,6 +1874,7 @@ def _probe_channel( backend = _val("wechat_backend", "wecom") if backend == "wechatmp": from ..channels.wechat.probe import validate_wechat_mp + return await validate_wechat_mp( _val("wechat_mp_app_id"), _val("wechat_mp_app_secret"), @@ -1598,6 +1882,7 @@ def _probe_channel( ) else: from ..channels.wechat.probe import validate_wecom + return await validate_wecom( _val("wechat_wecom_corp_id"), _val("wechat_wecom_secret"), @@ -1605,6 +1890,7 @@ def _probe_channel( ) elif ch_name == "feishu": from ..channels.feishu.probe import validate_feishu_credentials + return await validate_feishu_credentials( _val("feishu_app_id"), _val("feishu_app_secret"), @@ -1612,6 +1898,7 @@ def _probe_channel( ) elif ch_name == "dingtalk": from ..channels.dingtalk.probe import validate_dingtalk + return await validate_dingtalk( _val("dingtalk_client_id"), _val("dingtalk_client_secret"), @@ -1619,6 +1906,7 @@ def _probe_channel( ) elif ch_name == "email": from ..channels.email.probe import validate_email_imap + return await validate_email_imap( _val("email_imap_host"), int(_val("email_imap_port", "993")), @@ -1628,12 +1916,14 @@ def _probe_channel( ) elif ch_name == "qq": from ..channels.qq.probe import validate_qq + return await validate_qq( _val("qq_app_id"), _val("qq_app_secret"), ) elif ch_name == "signal": from ..channels.signal.probe import validate_signal + return await validate_signal( _val("signal_phone_number"), _val("signal_cli_path", "signal-cli"), @@ -1647,6 +1937,7 @@ def _probe_channel( loop = asyncio.get_event_loop() if loop.is_running(): import nest_asyncio # type: ignore[import-untyped] + nest_asyncio.apply() except RuntimeError: loop = asyncio.new_event_loop() @@ -1657,16 +1948,21 @@ def _probe_channel( console.print(f" [green]✓ {detail}[/green]") else: console.print(f" [yellow]⚠ {detail}[/yellow]") - console.print(" [dim]Channel will still be enabled — check credentials later.[/dim]") + console.print( + " [dim]Channel will still be enabled — check credentials later.[/dim]" + ) except Exception as e: console.print(f" [yellow]⚠ Could not validate: {e}[/yellow]") - console.print(" [dim]Channel will still be enabled — check credentials later.[/dim]") + console.print( + " [dim]Channel will still be enabled — check credentials later.[/dim]" + ) # ============================================================================= # Progress Rendering (for tests and potential future use) # ============================================================================= + def render_progress(current_step: int, completed: set[int]) -> Panel: """Render the progress indicator panel. @@ -1713,6 +2009,7 @@ def render_progress(current_step: int, completed: set[int]) -> Panel: # Main onboard function # ============================================================================= + def run_onboard(skip_validation: bool = False) -> bool: """Run the interactive onboarding wizard. @@ -1760,6 +2057,8 @@ def run_onboard(skip_validation: bool = False) -> bool: config.siliconflow_api_key = new_key elif provider == "openrouter": config.openrouter_api_key = new_key + elif provider in ("zhipu", "zhipu-code"): + config.zhipu_api_key = new_key elif provider == "custom": config.custom_api_key = new_key else: @@ -1775,6 +2074,8 @@ def run_onboard(skip_validation: bool = False) -> bool: current = config.siliconflow_api_key elif provider == "openrouter": current = config.openrouter_api_key + elif provider in ("zhipu", "zhipu-code"): + current = config.zhipu_api_key elif provider == "custom": current = config.custom_api_key else: @@ -1783,7 +2084,9 @@ def run_onboard(skip_validation: bool = False) -> bool: _print_step_skipped("API Key", "not set") # Step 3: Model - model = _step_model(config, provider, ollama_detected_models=ollama_detected_models) + model = _step_model( + config, provider, ollama_detected_models=ollama_detected_models + ) config.model = model # Step 4: Tavily Key diff --git a/EvoScientist/config/settings.py b/EvoScientist/config/settings.py index a618ea6..066ea50 100644 --- a/EvoScientist/config/settings.py +++ b/EvoScientist/config/settings.py @@ -19,6 +19,7 @@ import yaml # Configuration paths # ============================================================================= + def get_config_dir() -> Path: """Get the configuration directory path. @@ -39,6 +40,7 @@ def get_config_path() -> Path: # Configuration dataclass # ============================================================================= + @dataclass class EvoScientistConfig: """EvoScientist configuration settings. @@ -63,6 +65,7 @@ class EvoScientistConfig: google_api_key: str = "" siliconflow_api_key: str = "" openrouter_api_key: str = "" + zhipu_api_key: str = "" custom_api_key: str = "" custom_base_url: str = "" ollama_base_url: str = "" @@ -182,6 +185,7 @@ class EvoScientistConfig: # Config file operations # ============================================================================= + def load_config() -> EvoScientistConfig: """Load configuration from file. @@ -235,6 +239,7 @@ def reset_config() -> None: # Config value operations # ============================================================================= + def _coerce_value(value: Any, field_type: Any) -> Any: """Coerce a value to the expected field type. @@ -322,6 +327,7 @@ _ENV_MAPPINGS = { "google_api_key": "GOOGLE_API_KEY", "siliconflow_api_key": "SILICONFLOW_API_KEY", "openrouter_api_key": "OPENROUTER_API_KEY", + "zhipu_api_key": "ZHIPU_API_KEY", "custom_api_key": "CUSTOM_API_KEY", "custom_base_url": "CUSTOM_BASE_URL", "ollama_base_url": "OLLAMA_BASE_URL", @@ -332,7 +338,9 @@ _ENV_MAPPINGS = { } -def get_effective_config(cli_overrides: dict[str, Any] | None = None) -> EvoScientistConfig: +def get_effective_config( + cli_overrides: dict[str, Any] | None = None, +) -> EvoScientistConfig: """Get effective configuration by merging all sources. Priority (highest to lowest): @@ -355,7 +363,9 @@ def get_effective_config(cli_overrides: dict[str, Any] | None = None) -> EvoScie for config_key, env_key in _ENV_MAPPINGS.items(): env_value = os.environ.get(env_key) if env_value: - field_info = next(f for f in fields(EvoScientistConfig) if f.name == config_key) + field_info = next( + f for f in fields(EvoScientistConfig) if f.name == config_key + ) try: data[config_key] = _coerce_value(env_value, field_info.type) except (ValueError, TypeError): @@ -391,6 +401,8 @@ def apply_config_to_env(config: EvoScientistConfig) -> None: os.environ["SILICONFLOW_API_KEY"] = config.siliconflow_api_key if config.openrouter_api_key and not os.environ.get("OPENROUTER_API_KEY"): os.environ["OPENROUTER_API_KEY"] = config.openrouter_api_key + if config.zhipu_api_key and not os.environ.get("ZHIPU_API_KEY"): + os.environ["ZHIPU_API_KEY"] = config.zhipu_api_key if config.custom_api_key and not os.environ.get("CUSTOM_API_KEY"): os.environ["CUSTOM_API_KEY"] = config.custom_api_key if config.custom_base_url and not os.environ.get("CUSTOM_BASE_URL"): diff --git a/EvoScientist/llm/models.py b/EvoScientist/llm/models.py index 1732755..5dd2c45 100644 --- a/EvoScientist/llm/models.py +++ b/EvoScientist/llm/models.py @@ -2,7 +2,7 @@ This module provides a unified interface for creating chat model instances with support for multiple providers (Anthropic, OpenAI, Google GenAI, NVIDIA, -SiliconFlow, OpenRouter, Ollama, and custom OpenAI-compatible endpoints) and +SiliconFlow, OpenRouter, ZhipuAI, Ollama, and custom OpenAI-compatible endpoints) and convenient short names for common models. """ @@ -15,12 +15,16 @@ from langchain.chat_models import init_chat_model _SILICONFLOW_BASE_URL = "https://api.siliconflow.cn/v1" _OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1" +_ZHIPU_BASE_URL = "https://open.bigmodel.cn/api/paas/v4" +_ZHIPU_CODE_BASE_URL = "https://open.bigmodel.cn/api/coding/paas/v4" # Third-party providers routed through the OpenAI provider with a custom base_url. # Maps provider name → (base_url or None, env var for API key). _THIRD_PARTY_PROVIDERS: dict[str, tuple[str | None, str]] = { "siliconflow": (_SILICONFLOW_BASE_URL, "SILICONFLOW_API_KEY"), "openrouter": (_OPENROUTER_BASE_URL, "OPENROUTER_API_KEY"), + "zhipu": (_ZHIPU_BASE_URL, "ZHIPU_API_KEY"), + "zhipu-code": (_ZHIPU_CODE_BASE_URL, "ZHIPU_API_KEY"), "custom": (None, "CUSTOM_API_KEY"), # base_url from CUSTOM_BASE_URL env } @@ -73,6 +77,12 @@ _MODEL_ENTRIES: list[tuple[str, str, str]] = [ ("qwen3.5-122b", "qwen/qwen3.5-122b-a10b", "openrouter"), ("gemini-3-flash", "google/gemini-3-flash-preview", "openrouter"), ("claude-sonnet-4.6", "anthropic/claude-sonnet-4.6", "openrouter"), + # Zhipu (智谱) + ("glm-5", "glm-5", "zhipu"), + ("glm-4.7", "glm-4.7", "zhipu"), + # Zhipu Code Plan (智谱代码计划) + ("glm-5", "glm-5", "zhipu-code"), + ("glm-4.7", "glm-4.7", "zhipu-code"), ] # Public dict for simple lookups (last entry wins for duplicate names). diff --git a/README.md b/README.md index 9d06c42..06abcc8 100644 --- a/README.md +++ b/README.md @@ -82,7 +82,7 @@ Going beyond traditional human-in-the-loop systems, EvoScientist introduces an A ## 📦 Installation > [!TIP] -> Requires **Python 3.11+**. We recommend [**uv**](https://docs.astral.sh/uv/) or **conda** for dependency management and virtual environments. +> Requires **Python 3.11+** (**< 3.14**). We recommend [**uv**](https://docs.astral.sh/uv/) or **conda** for dependency management and virtual environments.
🪛 Install uv (if you don't have it) diff --git a/README.zh-CN.md b/README.zh-CN.md index d36c97f..48340d1 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -92,7 +92,7 @@ EvoScientist 超越了传统的人在回路(Human-in-the-Loop)模式,引 ## 📦 安装 > [!TIP] -> 需要 **Python 3.11+**。推荐使用 [**uv**](https://docs.astral.sh/uv/) 或 **conda** 进行依赖管理和虚拟环境管理。 +> 需要 **Python 3.11+**(**< 3.14**)。推荐使用 [**uv**](https://docs.astral.sh/uv/) 或 **conda** 进行依赖管理和虚拟环境管理。
🪛 安装 uv(如果尚未安装) diff --git a/architecture_design.md b/architecture_design.md new file mode 100644 index 0000000..7f2bf41 --- /dev/null +++ b/architecture_design.md @@ -0,0 +1,664 @@ +# Memory-Augmented LLM Architecture Design +## 类DeepSeek Engram存算分离架构方案 + +--- + +## 1. 核心设计理念 + +### 1.1 存算分离原则 +``` +┌─────────────────────────────────────────────────────────────┐ +│ GPU HBM (高速计算) │ +│ ┌─────────────┐ ┌─────────────┐ ┌─────────────┐ │ +│ │ Attention │ │ MoE │ │ FFN │ │ +│ │ Layers │ │ Experts │ │ (动态推理) │ │ +│ └─────────────┘ └─────────────┘ └─────────────┘ │ +│ ↑ ↑ │ +│ │ 检索增强 │ │ +│ └─────────────────┼─────────────────────────────────┤ +│ ↓ │ +│ ┌─────────────────────────────────────────────────────┐ │ +│ │ Memory Layer Interface │ │ +│ └─────────────────────────────────────────────────────┘ │ +└─────────────────────────────────────────────────────────────┘ + ↕ PCIe (异步预取) +┌─────────────────────────────────────────────────────────────┐ +│ CPU DRAM (大容量存储) │ +│ ┌─────────────────────────────────────────────────────┐ │ +│ │ Static Knowledge Memory Bank │ │ +│ │ (N-gram Hash → Embedding Table, 100B+ params) │ │ +│ │ │ │ +│ │ ┌──────────┐ ┌──────────┐ ┌──────────┐ │ │ +│ │ │ Unigram │ │ Bigram │ │ Trigram │ ... │ │ +│ │ │ Table │ │ Table │ │ Table │ │ │ +│ │ └──────────┘ └──────────┘ └──────────┘ │ │ +│ └─────────────────────────────────────────────────────┘ │ +└─────────────────────────────────────────────────────────────┘ +``` + +### 1.2 设计目标 +| 目标 | 指标 | 实现方式 | +|------|------|----------| +| 静态知识存储 | 100B+ params in DRAM | N-gram hash embedding | +| 推理速度 | <3% throughput penalty | 异步PCIe预取 | +| 知识召回 | +3-5 benchmark points | Multi-head hashing | +| 长上下文 | 128K+ tokens | Memory作为外部知识库 | + +--- + +## 2. 架构详细设计 + +### 2.1 整体模型结构 + +```python +# 架构层次示意 +MemoryAugmentedTransformer( + # Stage 1: Embedding + 早期层 + embedding: Embedding(vocab_size, hidden_dim), + early_layers: [ + TransformerBlock(with_memory=True), # Layer 2 插入 Memory + TransformerBlock(with_memory=False), + ], + + # Stage 2: 中间层 (MoE) + middle_layers: [ + MoEBlock(num_experts=64, top_k=4), + MoEBlock(num_experts=64, top_k=4), + # ... + ], + + # Stage 3: 后期层 + Memory + late_layers: [ + TransformerBlock(with_memory=False), + TransformerBlock(with_memory=True), # Layer 15 插入 Memory + ], + + # Memory Module + memory_module: ConditionalMemory( + ngram_orders=[1, 2, 3], + num_heads=8, + head_dim=1280, + offload_to_dram=True + ), + + # Output + lm_head: Linear(hidden_dim, vocab_size) +) +``` + +### 2.2 Conditional Memory Layer 设计 + +```python +class ConditionalMemoryLayer(nn.Module): + """ + 核心Memory Layer实现 + 基于 N-gram Hashing + Multi-Head Retrieval + Gating + """ + def __init__( + self, + hidden_dim: int = 4096, + ngram_orders: List[int] = [1, 2, 3], + num_heads: int = 8, + head_dim: int = 1280, + vocab_size: int = 128000, + offload_to_dram: bool = True + ): + super().__init__() + self.ngram_orders = ngram_orders + self.num_heads = num_heads + + # Multi-head hash projections + self.hash_projections = nn.ModuleList([ + nn.Linear(hidden_dim, head_dim) + for _ in range(num_heads) + ]) + + # Gating mechanism (context-aware) + self.gate = nn.Sequential( + nn.Linear(hidden_dim, hidden_dim // 4), + nn.SiLU(), + nn.Linear(hidden_dim // 4, head_dim), + nn.Sigmoid() + ) + + # Output projection + self.output_proj = nn.Linear(head_dim, hidden_dim) + + # Memory Bank (可offload到DRAM) + self.memory_bank = NgramMemoryBank( + vocab_size=vocab_size, + ngram_orders=ngram_orders, + num_heads=num_heads, + head_dim=head_dim, + offload_to_dram=offload_to_dram + ) + + def forward(self, hidden_states: torch.Tensor, input_ids: torch.Tensor): + """ + Args: + hidden_states: [batch, seq_len, hidden_dim] + input_ids: [batch, seq_len] + Returns: + memory_output: [batch, seq_len, hidden_dim] + """ + batch_size, seq_len, _ = hidden_states.shape + + # 1. Compute N-gram hashes for each position + ngram_hashes = self._compute_ngram_hashes(input_ids) # [batch, seq_len, num_orders] + + # 2. Multi-head retrieval from memory bank + retrieved = self.memory_bank.retrieve(ngram_hashes) # [batch, seq_len, num_heads, head_dim] + + # 3. Aggregate across heads + aggregated = retrieved.mean(dim=2) # [batch, seq_len, head_dim] + + # 4. Context-aware gating + gate_weights = self.gate(hidden_states) # [batch, seq_len, head_dim] + gated_memory = gate_weights * aggregated + + # 5. Project back to hidden dim + output = self.output_proj(gated_memory) + + return output + + def _compute_ngram_hashes(self, input_ids: torch.Tensor) -> torch.Tensor: + """ + 计算多阶N-gram的hash索引 + 使用 rolling hash 提高效率 + """ + # 简化实现示意 + hashes = [] + for order in self.ngram_orders: + # 使用多个独立的hash函数 + for head in range(self.num_heads): + hash_val = self._rolling_hash(input_ids, order, seed=head) + hashes.append(hash_val) + return torch.stack(hashes, dim=-1) +``` + +### 2.3 N-gram Memory Bank (支持DRAM Offloading) + +```python +class NgramMemoryBank(nn.Module): + """ + 大规模N-gram Embedding存储 + 支持DRAM offloading + 异步预取 + """ + def __init__( + self, + vocab_size: int, + ngram_orders: List[int], + num_heads: int, + head_dim: int, + offload_to_dram: bool = True + ): + super().__init__() + self.offload_to_dram = offload_to_dram + self.head_dim = head_dim + + # 计算每个n-gram order的table大小 + # 使用Product Quantization压缩 + self.pq_compressor = ProductQuantizer( + vector_dim=head_dim, + num_subvectors=8, # 8个子空间 + bits_per_subvector=8 # 每个子空间256个centroid + ) + + # Hash tables for each n-gram order and head + # 存储PQ codes而非完整向量 + self.hash_tables = nn.ParameterDict() + for order in ngram_orders: + table_size = self._estimate_table_size(vocab_size, order) + # PQ codes: [table_size, num_subvectors] + self.hash_tables[f'n{order}'] = nn.Parameter( + torch.zeros(table_size, 8, dtype=torch.uint8), + requires_grad=False + ) + + # Centroids for PQ (这些需要学习) + self.centroids = nn.Parameter( + torch.randn(256, head_dim // 8) # 256 centroids per subspace + ) + + if offload_to_dram: + self._setup_dram_offloading() + + def retrieve(self, ngram_hashes: torch.Tensor) -> torch.Tensor: + """ + 根据hash检索embedding + 支持batch retrieval + """ + batch_size, seq_len, num_hashes = ngram_hashes.shape + + # 异步预取到GPU + if self.offload_to_dram: + self._async_prefetch(ngram_hashes) + + # 从hash table获取PQ codes + pq_codes = self._lookup_pq_codes(ngram_hashes) # [batch, seq_len, num_heads, num_subvectors] + + # 反量化重建向量 + reconstructed = self.pq_compressor.decode(pq_codes, self.centroids) + + return reconstructed + + def _async_prefetch(self, hashes: torch.Tensor): + """ + 异步PCIe预取,隐藏延迟 + """ + # 预测接下来需要的hash + prefetch_hashes = self._predict_next_accesses(hashes) + + # 异步拷贝到GPU pinned memory + torch.cuda.current_stream().synchronize() + with torch.cuda.stream(self.prefetch_stream): + self._copy_to_gpu_async(prefetch_hashes) +``` + +### 2.4 Tokenizer Compression (减少词汇表冗余) + +```python +class CompressedTokenizer: + """ + NFKC规范化 + 大小写折叠 + 减少23%词汇表大小 + """ + def __init__(self, base_tokenizer): + self.base = base_tokenizer + self.canonical_map = self._build_canonical_map() + + def normalize(self, text: str) -> str: + """ + NFKC → NFD → strip accents → lowercase → whitespace collapse + """ + import unicodedata + + # NFKC normalization + text = unicodedata.normalize('NFKC', text) + # NFD to separate accents + text = unicodedata.normalize('NFD', text) + # Strip combining marks (accents) + text = ''.join(c for c in text if unicodedata.category(c) != 'Mn') + # Lowercase + text = text.lower() + # Collapse whitespace + text = ' '.join(text.split()) + + return text + + def tokenize(self, text: str) -> List[int]: + normalized = self.normalize(text) + return self.base.encode(normalized) +``` + +--- + +## 3. 预训练策略 + +### 3.1 两阶段预训练 + +``` +Stage 1: Memory-agnostic Pretraining (0-70% training) +──────────────────────────────────────────────────── +目标: 学习基础语言表示 +配置: + - Memory layer: 冻结或随机初始化 + - 主干网络: 正常训练 + - 数据: 通用文本语料 + +Stage 2: Memory-aware Pretraining (70-100% training) +──────────────────────────────────────────────────── +目标: 学习将静态知识写入memory +配置: + - Memory layer: 解冻,开始学习 + - Memory更新策略: + * 频繁出现的n-gram → 强化记忆 + * 罕见n-gram → 弱化或忽略 + - 数据: 高质量知识密集型语料 (Wikipedia, 教科书等) +``` + +### 3.2 Memory学习目标 + +```python +class MemoryLearningObjective(nn.Module): + """ + 训练Memory存储有用的静态知识 + """ + def __init__(self, temperature: float = 0.1): + self.temperature = temperature + + def compute_loss( + self, + hidden_states: torch.Tensor, # 当前层表示 + memory_output: torch.Tensor, # Memory检索结果 + target_hidden: torch.Tensor, # 下一层目标表示 + next_token_logits: torch.Tensor # 语言模型预测 + ): + # Loss 1: Memory应该提供有用的信息 + # 如果memory_output能预测target,说明存储了正确知识 + memory_pred = F.linear(memory_output, self.pred_head.weight) + knowledge_loss = F.mse_loss(memory_pred, target_hidden.detach()) + + # Loss 2: Gating应该学会选择性使用memory + # 动态内容 → gate close + # 静态内容 → gate open + gate_entropy = -(gate_weights * torch.log(gate_weights + 1e-8)).sum(-1).mean() + + # Loss 3: 稀疏性 - 不要对所有内容都使用memory + sparsity_loss = gate_weights.mean() + + return knowledge_loss + 0.1 * gate_entropy + 0.01 * sparsity_loss +``` + +### 3.3 静态 vs 动态知识分类 + +```python +class KnowledgeClassifier: + """ + 判断哪些知识应该存入memory + """ + def __init__(self): + # 静态知识特征 + self.static_patterns = [ + r'\d{4}-\d{2}-\d{2}', # 日期格式 + r'[A-Z][a-z]+ [A-Z][a-z]+', # 人名 + r'\b[A-Z]{2,}\b', # 缩写 + # ... 更多模式 + ] + + def should_memorize(self, text: str, frequency: int) -> float: + """ + 返回应该记忆的置信度 [0, 1] + + 高频 + 静态模式 → 高置信度 + 低频 + 动态内容 → 低置信度 + """ + score = 0.0 + + # 频率因素 (高频更重要) + freq_score = min(frequency / 10000, 1.0) + + # 模式匹配 + pattern_score = 0.0 + for pattern in self.static_patterns: + if re.search(pattern, text): + pattern_score = 1.0 + break + + return 0.7 * freq_score + 0.3 * pattern_score +``` + +--- + +## 4. 推理时检索机制 + +### 4.1 推理流程 + +``` +Input Text + │ + ▼ +┌─────────────────┐ +│ Tokenizer │ ← 规范化处理 +└────────┬────────┘ + │ + ▼ +┌─────────────────┐ +│ Embedding │ +└────────┬────────┘ + │ + ▼ +┌─────────────────┐ ┌──────────────────┐ +│ Layer 1 │ │ │ +└────────┬────────┘ │ │ + │ │ │ + ▼ │ │ +┌─────────────────┐ │ DRAM Memory │ +│ Layer 2 + MEM │◄────┤ Bank (100B) │ +└────────┬────────┘ │ │ + │ │ N-gram Hash │ + ▼ │ Tables │ +┌─────────────────┐ │ │ +│ Layers 3-14 │ │ │ +│ (MoE) │ └──────────────────┘ +└────────┬────────┘ ▲ + │ │ + ▼ │ +┌─────────────────┐ │ +│ Layer 15 + MEM │──────────────┘ +└────────┬────────┘ + │ + ▼ +┌─────────────────┐ +│ Output Layers │ +└────────┬────────┘ + │ + ▼ + Predictions +``` + +### 4.2 推理优化策略 + +```python +class InferenceOptimizer: + """ + 推理时优化 + """ + def __init__(self, model): + self.model = model + self.kv_cache = {} + self.memory_cache = LRUCache(size=10000) + + @torch.no_grad() + def generate( + self, + input_ids: torch.Tensor, + max_new_tokens: int = 100 + ): + # 1. Prefill阶段: 批量处理prompt + past_key_values = None + + for layer in self.model.layers: + # 检查memory cache + ngram_hashes = self._compute_hashes(input_ids) + + # Cache hit → 直接使用 + # Cache miss → 从DRAM获取 + memory_out = self._retrieve_with_cache( + ngram_hashes, + layer.memory_layer + ) + + # 2. Decode阶段: 增量生成 + for _ in range(max_new_tokens): + # 只处理最后一个token + # 利用KV cache避免重复计算 + pass + + def _retrieve_with_cache(self, hashes, memory_layer): + """ + LRU缓存 + 预取 + """ + # 检查缓存 + cache_key = hashes.cpu().tolist() + if cache_key in self.memory_cache: + return self.memory_cache[cache_key] + + # 从DRAM获取 + result = memory_layer.retrieve(hashes) + + # 更新缓存 + self.memory_cache[cache_key] = result + + return result +``` + +### 4.3 动态知识处理 + +```python +class DynamicKnowledgeHandler: + """ + 处理动态变化的knowledge (news, 实时数据等) + """ + def __init__(self, model): + self.model = model + self.external_retriever = None # 可接RAG系统 + + def process_dynamic_query(self, query: str, context: str): + """ + 动态知识 → 不使用memory,走attention路径 + 静态知识 → 使用memory,减轻attention负担 + """ + # 1. 分类query类型 + is_dynamic = self._is_dynamic_query(query) + + if is_dynamic: + # 动态查询: 关闭memory gate,依赖attention + with torch.no_grad(): + self.model.set_memory_gate_bias(-10.0) # 强制关闭 + output = self.model(query, context) + self.model.reset_memory_gate_bias() + else: + # 静态查询: 正常使用memory + output = self.model(query) + + return output + + def _is_dynamic_query(self, query: str) -> bool: + """ + 判断是否需要动态知识 + """ + dynamic_keywords = ['最新', '今天', '近期', 'current', 'latest', 'now'] + return any(kw in query.lower() for kw in dynamic_keywords) +``` + +--- + +## 5. 实验验证计划 + +### 5.1 实验阶段 + +| 阶段 | 目标 | 数据集 | 成功指标 | +|------|------|--------|----------| +| **Stage 1: Baseline** | 建立无memory基线 | Pile | Loss curve, PPL | +| **Stage 2: Memory Integration** | 集成memory layer | Pile + Wiki | +2-3 MMLU | +| **Stage 3: Offloading** | 测试DRAM offloading | - | <3% throughput loss | +| **Stage 4: Ablation** | 消融各组件 | MMLU, BBH | 各组件贡献 | +| **Stage 5: Scaling** | 扩展到更大模型 | 全量数据 | 对标Engram-27B | + +### 5.2 消融实验设计 + +```python +ablation_configs = [ + # 配置1: 无memory (baseline) + {'memory_enabled': False}, + + # 配置2: 单层memory + {'memory_enabled': True, 'memory_layers': [2]}, + + # 配置3: 双层memory (完整) + {'memory_enabled': True, 'memory_layers': [2, 15]}, + + # 配置4: 无multi-head hashing + {'memory_enabled': True, 'num_hash_heads': 1}, + + # 配置5: 无gating + {'memory_enabled': True, 'use_gating': False}, + + # 配置6: 无tokenizer压缩 + {'memory_enabled': True, 'compress_tokenizer': False}, + + # 配置7: PQ压缩 + {'memory_enabled': True, 'use_pq': True, 'pq_bits': 8}, +] +``` + +### 5.3 评估基准 + +```python +evaluation_suite = { + # 知识密集型任务 (memory应该帮助) + 'knowledge': [ + 'MMLU', # 多领域知识 + 'TriviaQA', # 事实问答 + 'NaturalQuestions', + ], + + # 推理任务 (主要依赖attention) + 'reasoning': [ + 'BBH', # 综合推理 + 'GSM8K', # 数学 + 'HumanEval', # 代码 + ], + + # 长上下文任务 + 'long_context': [ + 'NeedleInAHaystack', # 检索能力 + 'LongBench', # 长文本理解 + ], + + # 效率指标 + 'efficiency': [ + 'throughput_tokens_per_sec', + 'memory_usage_gb', + 'latency_ms', + ], +} +``` + +--- + +## 6. 实现路线图 + +### Phase 1: 核心组件 (2-3周) +- [ ] N-gram hash函数实现 +- [ ] Memory bank基础结构 +- [ ] Multi-head retrieval +- [ ] Gating mechanism + +### Phase 2: 集成训练 (3-4周) +- [ ] 与Transformer集成 +- [ ] 两阶段预训练pipeline +- [ ] Memory学习目标 +- [ ] 小规模实验 (1B model) + +### Phase 3: 优化部署 (2-3周) +- [ ] DRAM offloading +- [ ] 异步预取 +- [ ] Product Quantization压缩 +- [ ] 推理优化 + +### Phase 4: 扩展验证 (3-4周) +- [ ] 中等规模实验 (7B model) +- [ ] 完整评估 +- [ ] 消融实验 +- [ ] 论文撰写 + +--- + +## 7. 参考资源 + +### 核心论文 +1. **DeepSeek Engram**: https://arxiv.org/abs/2601.07372 +2. **RETRO**: "Improving language models by retrieving from trillions of tokens" +3. **Neural Turing Machines**: Graves et al., 2014 +4. **Product Quantization**: Jégou et al., 2011 + +### 开源实现 +- Engram GitHub: https://github.com/deepseek-ai/Engram +- RETRO: https://nn.labml.ai/transformers/retro/ +- FAISS (PQ): https://github.com/facebookresearch/faiss + +--- + +## 8. 下一步行动 + +1. **立即可做**: 搭建基础Transformer backbone + 简单memory layer +2. **短期目标**: 在小数据集验证memory学习有效性 +3. **中期目标**: 完整两阶段预训练实验 +4. **长期目标**: 扩展到10B+规模,对标Engram + +需要我开始实现哪个部分?我建议从以下开始: +- A) 核心Memory Layer实现 (PyTorch) +- B) 训练pipeline搭建 +- C) 小规模验证实验设计 diff --git a/tests/test_llm.py b/tests/test_llm.py index dbe767a..b2dda52 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -35,7 +35,7 @@ class TestModelsRegistry: def test_entries_are_valid_tuples(self): """Test that _MODEL_ENTRIES contains valid (name, model_id, provider) tuples.""" - valid_providers = {"anthropic", "openai", "google-genai", "nvidia", "siliconflow", "openrouter"} + valid_providers = {"anthropic", "openai", "google-genai", "nvidia", "siliconflow", "openrouter", "zhipu", "zhipu-code"} for entry in _MODEL_ENTRIES: assert len(entry) == 3, f"Entry {entry} doesn't have 3 elements" name, model_id, provider = entry