Text Generation
Transformers
Safetensors
GGUF
English
causal-lm
qwen2.5
reasoning
code-generation
Mixture of Experts
qlora
multimodal
tool-use
Eval Results (legacy)
conversational
Instructions to use ram1234598766/Cesium2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use ram1234598766/Cesium2 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="ram1234598766/Cesium2") messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("ram1234598766/Cesium2", device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- llama.cpp
How to use ram1234598766/Cesium2 with llama.cpp:
Install (macOS, Linux)
curl -LsSf https://llama.app/install.sh | sh # Start a local OpenAI-compatible server with a web UI: llama serve -hf ram1234598766/Cesium2:Q8_0 # Run inference directly in the terminal: llama cli -hf ram1234598766/Cesium2:Q8_0
Install from WinGet (Windows)
winget install llama.cpp # Start a local OpenAI-compatible server with a web UI: llama serve -hf ram1234598766/Cesium2:Q8_0 # Run inference directly in the terminal: llama cli -hf ram1234598766/Cesium2:Q8_0
Use pre-built binary
# Download pre-built binary from: # https://github.com/ggerganov/llama.cpp/releases # Start a local OpenAI-compatible server with a web UI: ./llama-server -hf ram1234598766/Cesium2:Q8_0 # Run inference directly in the terminal: ./llama-cli -hf ram1234598766/Cesium2:Q8_0
Build from source code
git clone https://github.com/ggerganov/llama.cpp.git cd llama.cpp cmake -B build cmake --build build -j --target llama-server llama-cli # Start a local OpenAI-compatible server with a web UI: ./build/bin/llama-server -hf ram1234598766/Cesium2:Q8_0 # Run inference directly in the terminal: ./build/bin/llama-cli -hf ram1234598766/Cesium2:Q8_0
Use Docker
docker model run hf.co/ram1234598766/Cesium2:Q8_0
- LM Studio
- Jan
- vLLM
How to use ram1234598766/Cesium2 with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "ram1234598766/Cesium2" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "ram1234598766/Cesium2", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/ram1234598766/Cesium2:Q8_0
- SGLang
How to use ram1234598766/Cesium2 with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "ram1234598766/Cesium2" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "ram1234598766/Cesium2", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "ram1234598766/Cesium2" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "ram1234598766/Cesium2", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Ollama
How to use ram1234598766/Cesium2 with Ollama:
ollama run hf.co/ram1234598766/Cesium2:Q8_0
- Unsloth Studio
How to use ram1234598766/Cesium2 with Unsloth Studio:
Install Unsloth Studio (macOS, Linux, WSL)
curl -fsSL https://unsloth.ai/install.sh | sh # Run unsloth studio unsloth studio -H 0.0.0.0 -p 8888 # Then open http://localhost:8888 in your browser # Search for ram1234598766/Cesium2 to start chatting
Install Unsloth Studio (Windows)
irm https://unsloth.ai/install.ps1 | iex # Run unsloth studio unsloth studio -H 0.0.0.0 -p 8888 # Then open http://localhost:8888 in your browser # Search for ram1234598766/Cesium2 to start chatting
Using HuggingFace Spaces for Unsloth
# No setup required # Open https://huggingface.co/spaces/unsloth/studio in your browser # Search for ram1234598766/Cesium2 to start chatting
- Pi
How to use ram1234598766/Cesium2 with Pi:
Start the llama.cpp server
# Install llama.cpp: brew install llama.cpp # Start a local OpenAI-compatible server: llama serve -hf ram1234598766/Cesium2:Q8_0
Configure the model in Pi
# Install Pi: npm install -g @mariozechner/pi-coding-agent # Add to ~/.pi/agent/models.json: { "providers": { "llama-cpp": { "baseUrl": "http://localhost:8080/v1", "api": "openai-completions", "apiKey": "none", "models": [ { "id": "ram1234598766/Cesium2:Q8_0" } ] } } }Run Pi
# Start Pi in your project directory: pi
- Docker Model Runner
How to use ram1234598766/Cesium2 with Docker Model Runner:
docker model run hf.co/ram1234598766/Cesium2:Q8_0
- Lemonade
How to use ram1234598766/Cesium2 with Lemonade:
Pull the model
# Download Lemonade from https://lemonade-server.ai/ lemonade pull ram1234598766/Cesium2:Q8_0
Run and chat with the model
lemonade run user.Cesium2-Q8_0
List all available models
lemonade list
- Hermes Agent
How to use ram1234598766/Cesium2 with Hermes Agent:
Start the llama.cpp server
# Install llama.cpp: brew install llama.cpp # Start a local OpenAI-compatible server: llama serve -hf ram1234598766/Cesium2:Q8_0
Configure Hermes
# Install Hermes: curl -fsSL https://hermes-agent.nousresearch.com/install.sh | bash hermes setup # Point Hermes at the local server: hermes config set model.provider custom hermes config set model.base_url http://127.0.0.1:8080/v1 hermes config set model.default ram1234598766/Cesium2:Q8_0
Run Hermes
hermes
- Atomic Chat
- OpenClaw
How to use ram1234598766/Cesium2 with OpenClaw:
Start the llama.cpp server
# Install llama.cpp: brew install llama.cpp # Start a local OpenAI-compatible server: llama serve -hf ram1234598766/Cesium2:Q8_0
Configure OpenClaw
# Install OpenClaw: npm install -g openclaw@latest # Register the local server and set it as the default model: openclaw onboard --non-interactive --mode local \ --auth-choice custom-api-key \ --custom-base-url http://127.0.0.1:8080/v1 \ --custom-model-id "ram1234598766/Cesium2:Q8_0" \ --custom-provider-id llama-cpp \ --custom-compatibility openai \ --custom-text-input \ --accept-risk \ --skip-health
Run OpenClaw
openclaw agent --local --agent main --message "Hello from Hugging Face"
| """ | |
| MORPH-AI Architecture v6 | |
| Advanced, memory-efficient architecture for local + mobile inference, engineered for | |
| excellent reasoning, code generation, and ALL capabilities. | |
| Core design: "System-1 / System-2" dual-path with dynamic compute. | |
| A Coordinator decides, per input, how much thinking to spend and which subsystems | |
| to activate. Memory-efficient attention (SDPA/Flash Attention) + Mixture of Depths | |
| (MoD) for dynamic layer skipping. Reasoning is an iterative, weight-tied refinement | |
| loop (System 2) over a compressed state space. Code structure is injected as a | |
| learned bias. A persistent scratchpad carries reasoning state across turns. | |
| v6 NEW: | |
| - Memory-efficient SDPA attention (PyTorch 2.0+ native, fallback-safe) | |
| - Mixture of Depths (MoD): dynamically skip transformer layers per token | |
| - Dynamic MoE with expert pruning and load-balanced routing | |
| - KV Cache quantization (INT8/INT4) for long-context memory efficiency | |
| - Multimodal Fusion Layer (text + vision + audio + video embeddings) | |
| - Tool Use Module (JSON-structured function calling with validation) | |
| - Document Understanding Module (PDF/DOCX/OCR with layout-aware parsing) | |
| - Video Understanding Module (temporal frame sampling + motion features) | |
| - Code Execution Sandbox (safe Python execution with AST validation) | |
| - Speculative Decoding support (draft + verification chain) | |
| - RoPE scaling for extended context windows | |
| - 8-bit optimizer compatibility + paged AdamW | |
| Modules: | |
| 1. Coordinator - routes between subsystems, predicts reasoning depth | |
| 2. MultiStepReasoner - iterative (System-2) refinement loop, weight-tied | |
| 3. CodeAwareBias - injects code structure (indent, brackets) as bias | |
| 4. ScratchpadMemory - persistent cross-turn working memory | |
| 5. VerifierHead - scores generations for best-of-n self-critique | |
| 6. MoEBlock - sparse top-k experts + load-balance loss + pruning | |
| 7. MemoryModule - persistent key-value memory (attention read) + quantization | |
| 8. SkillTokenModule - hot-swappable skill embeddings | |
| 9. DepthEmbeddings - predicts task depth, injects conditioning vector | |
| 10. MixtureOfDepths - dynamically skip transformer layers per token | |
| 11. MultimodalFusion - fuse text + vision + audio + video embeddings | |
| 12. ToolUseModule - JSON-structured function calling with validation | |
| 13. DocumentModule - PDF/DOCX/OCR with layout-aware parsing | |
| 14. VideoModule - temporal frame sampling + motion features | |
| 15. CodeSandbox - safe Python execution with AST validation | |
| Trainable end-to-end with 4-bit QLoRA on a free Colab T4 (~16GB). | |
| """ | |
| import math | |
| from dataclasses import dataclass | |
| from typing import Dict, List, Optional, Tuple | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from peft import LoraConfig, TaskType, get_peft_model | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| from transformers.modeling_outputs import ModelOutput | |
| class MorphConfig: | |
| base_model: str = "Qwen/Qwen2.5-1.5B-Instruct" | |
| # skills | |
| num_skill_tokens: int = 64 | |
| # adaptive compute | |
| max_depth: int = 12 | |
| coordinator_hidden: int = 256 | |
| adaptive_threshold: float = 0.5 | |
| # MoD - Mixture of Depths | |
| use_mod: bool = True | |
| mod_hidden: int = 128 | |
| mod_dropout: float = 0.1 | |
| mod_keep_prob: float = 0.8 | |
| # memory-efficient attention | |
| use_sdpa: bool = True | |
| attn_dropout: float = 0.0 | |
| # LoRA | |
| lora_rank: int = 16 | |
| lora_alpha: int = 32 | |
| lora_dropout: float = 0.05 | |
| # MoE | |
| num_experts: int = 4 | |
| max_experts: int = 64 | |
| expert_hidden: int = 512 | |
| moe_top_k: int = 2 | |
| moe_aux_weight: float = 0.01 | |
| moe_prune_threshold: float = 0.02 | |
| moe_expand_threshold: float = 0.15 | |
| # MoD - Mixture of Depths | |
| use_mod: bool = True | |
| mod_hidden: int = 128 | |
| mod_dropout: float = 0.1 | |
| mod_keep_prob: float = 0.8 | |
| mod_temperature: float = 1.0 | |
| mod_temperature_anneal: float = 0.995 | |
| # multi-head CoT reasoning | |
| num_cot_heads: int = 4 | |
| cot_hidden: int = 256 | |
| # memory | |
| memory_size: int = 1024 | |
| memory_dim: int = 768 | |
| memory_quantize: bool = True | |
| memory_quant_bits: int = 8 | |
| # System-2 reasoning loop | |
| reasoner_dim: int = 512 | |
| reasoner_heads: int = 4 | |
| reasoner_ff: int = 768 | |
| max_steps: int = 4 | |
| # scratchpad | |
| scratch_dim: int = 512 | |
| # code awareness | |
| code_feat_dim: int = 7 | |
| # verifier | |
| verifier_weight: float = 0.05 | |
| # multimodal | |
| vision_dim: int = 768 | |
| audio_dim: int = 768 | |
| video_dim: int = 768 | |
| fusion_hidden: int = 512 | |
| # tool use | |
| max_tools: int = 16 | |
| tool_hidden: int = 256 | |
| # document | |
| doc_max_pages: int = 10 | |
| doc_hidden: int = 256 | |
| # video | |
| video_max_frames: int = 8 | |
| video_hidden: int = 256 | |
| # code sandbox | |
| sandbox_timeout: float = 5.0 | |
| sandbox_max_memory: int = 128 # MB | |
| # speculative decoding | |
| use_speculative: bool = False | |
| draft_layers: int = 2 | |
| # RoPE scaling for extended context | |
| rope_scaling: Optional[dict] = None | |
| # plugin architecture | |
| plugin_dir: Optional[str] = None | |
| # training | |
| max_seq_len: int = 8192 | |
| # quantization | |
| load_in_8bit: bool = False | |
| load_in_4bit: bool = True | |
| bnb_4bit_compute_dtype: str = "bfloat16" | |
| bnb_4bit_quant_type: str = "nf4" | |
| bnb_4bit_use_double_quant: bool = True | |
| # --------------------------------------------------------------------------- | |
| # Expert + MoE | |
| # --------------------------------------------------------------------------- | |
| class Expert(nn.Module): | |
| """Single MoE expert - lightweight SiLU FFN.""" | |
| def __init__(self, hidden_dim: int, expert_hidden: int): | |
| super().__init__() | |
| self.w1 = nn.Linear(hidden_dim, expert_hidden, bias=False) | |
| self.w2 = nn.Linear(expert_hidden, hidden_dim, bias=False) | |
| self.act = nn.SiLU() | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.w2(self.act(self.w1(x))) | |
| class MoEBlock(nn.Module): | |
| """ | |
| Sparse Mixture of Experts. Only top-k experts activate per token, giving | |
| k*expert_hidden capacity for ~k/E of the FFN compute. Returns routed output | |
| plus a load-balancing auxiliary loss. | |
| """ | |
| def __init__(self, hidden_dim: int, num_experts: int, expert_hidden: int, top_k: int): | |
| super().__init__() | |
| self.num_experts = num_experts | |
| self.top_k = top_k | |
| self.gate = nn.Linear(hidden_dim, num_experts, bias=False) | |
| self.experts = nn.ModuleList([ | |
| Expert(hidden_dim, expert_hidden) for _ in range(num_experts) | |
| ]) | |
| def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: | |
| B, T, H = x.shape | |
| flat = x.reshape(-1, H) | |
| gate_logits = self.gate(flat) # (B*T, E) | |
| probs = F.softmax(gate_logits, dim=-1) | |
| topk_vals, topk_idx = torch.topk(gate_logits, self.top_k, dim=-1) | |
| topk_vals = F.softmax(topk_vals, dim=-1) | |
| routing = torch.zeros_like(probs) | |
| routing.scatter_(1, topk_idx, topk_vals) | |
| out = torch.zeros_like(flat) | |
| for i, expert in enumerate(self.experts): | |
| sel = routing[:, i] > 0 | |
| if sel.any(): | |
| out[sel] += routing[sel, i].unsqueeze(-1) * expert(flat[sel]) | |
| f_i = routing.mean(0) | |
| P_i = probs.mean(0) | |
| aux = (f_i * P_i).sum() * self.num_experts | |
| return out.view(B, T, H), aux | |
| # --------------------------------------------------------------------------- | |
| # Persistent key-value memory | |
| # --------------------------------------------------------------------------- | |
| class MemoryModule(nn.Module): | |
| """Persistent key-value memory. Differentiable attention read; EMA write.""" | |
| def __init__(self, memory_size: int, memory_dim: int, hidden_dim: int): | |
| super().__init__() | |
| self.memory_size = memory_size | |
| self.memory_dim = memory_dim | |
| self.key_proj = nn.Linear(hidden_dim, memory_dim) | |
| self.query_proj = nn.Linear(hidden_dim, memory_dim) | |
| self.val_proj = nn.Linear(hidden_dim, memory_dim) | |
| self.out_proj = nn.Linear(memory_dim, hidden_dim) | |
| self.mem_k = nn.Parameter(torch.randn(memory_size, memory_dim) * 0.02) | |
| self.mem_v = nn.Parameter(torch.randn(memory_size, memory_dim) * 0.02) | |
| self.mem_k_buf = None | |
| self.mem_v_buf = None | |
| def read(self, hidden: torch.Tensor) -> torch.Tensor: | |
| query = self.query_proj(hidden) # (B, T, D) | |
| keys = self.mem_k_buf if self.mem_k_buf is not None else self.mem_k | |
| vals = self.mem_v_buf if self.mem_v_buf is not None else self.mem_v | |
| attn = torch.matmul(query, keys.T) | |
| attn = F.softmax(attn / math.sqrt(self.memory_dim), dim=-1) | |
| retrieved = torch.matmul(attn, vals) | |
| return self.out_proj(retrieved) | |
| def write(self, hidden: torch.Tensor): | |
| with torch.no_grad(): | |
| key = self.key_proj(hidden).mean(1) | |
| val = self.val_proj(hidden).mean(1) | |
| if self.mem_k_buf is None: | |
| self.mem_k_buf = self.mem_k.detach().clone() | |
| self.mem_v_buf = self.mem_v.detach().clone() | |
| for k, v in zip(key, val): | |
| if self.mem_k_buf.is_cuda and k.is_cpu: | |
| k = k.cuda() | |
| if self.mem_v_buf.is_cuda and v.is_cpu: | |
| v = v.cuda() | |
| n = min(k.size(0), self.memory_size) | |
| alpha = 0.1 | |
| self.mem_k_buf[:n] = (1 - alpha) * self.mem_k_buf[:n] + alpha * k[:n] | |
| self.mem_v_buf[:n] = (1 - alpha) * self.mem_v_buf[:n] + alpha * v[:n] | |
| # --------------------------------------------------------------------------- | |
| # Skill + depth | |
| # --------------------------------------------------------------------------- | |
| class SkillTokenModule(nn.Module): | |
| """Hot-swappable skill embeddings injected into the input embedding stream.""" | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| self.num_skill_tokens = config.num_skill_tokens | |
| self.skill_embeddings = nn.Embedding(config.num_skill_tokens, hidden_dim) | |
| nn.init.normal_(self.skill_embeddings.weight, std=0.02) | |
| self.skill_proj = nn.Linear(hidden_dim, hidden_dim) | |
| nn.init.zeros_(self.skill_proj.weight) | |
| nn.init.zeros_(self.skill_proj.bias) | |
| def forward(self, skill_indices: Optional[torch.Tensor]) -> Optional[torch.Tensor]: | |
| if skill_indices is None or skill_indices.numel() == 0: | |
| return None | |
| emb = self.skill_embeddings(skill_indices) | |
| return self.skill_proj(emb.mean(0, keepdim=True)) # (1, H) | |
| class DepthEmbeddings(nn.Module): | |
| """Predicts task depth from the final hidden state and injects a conditioning vector.""" | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| self.max_depth = config.max_depth | |
| self.depth_embeddings = nn.Embedding(config.max_depth + 1, hidden_dim) | |
| nn.init.normal_(self.depth_embeddings.weight, std=0.02) | |
| self.depth_predictor = nn.Sequential( | |
| nn.Linear(hidden_dim, 128), | |
| nn.GELU(), | |
| nn.Linear(128, config.max_depth + 1), | |
| nn.Softmax(dim=-1), | |
| ) | |
| def forward(self, hidden: torch.Tensor, force_depth: Optional[int] = None) -> Tuple[torch.Tensor, torch.Tensor]: | |
| last = hidden[:, -1, :] | |
| dist = self.depth_predictor(last) # (B, max_depth+1) | |
| if force_depth is not None: | |
| d = torch.clamp(torch.tensor(force_depth, device=hidden.device).long(), 0, self.max_depth) | |
| emb = self.depth_embeddings(d).unsqueeze(0) | |
| dist = F.one_hot(d, num_classes=self.max_depth + 1).float() | |
| else: | |
| # differentiable soft mixture of depth embeddings: trains the | |
| # depth_predictor + depth_embeddings end-to-end through the logits | |
| emb = dist @ self.depth_embeddings.weight # (B, H) | |
| return emb, dist | |
| # --------------------------------------------------------------------------- | |
| # Coordinator (System-1/System-2 controller) | |
| # --------------------------------------------------------------------------- | |
| class Coordinator(nn.Module): | |
| """ | |
| Hierarchical controller. Given the base hidden state, decides: | |
| gates = [think, code, memory, scratch] (per-sequence, in [0,1]) | |
| steps = number of System-2 refinement iterations (0..max_steps) | |
| """ | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| self.max_steps = config.max_steps | |
| self.hidden = nn.Sequential( | |
| nn.Linear(hidden_dim, config.coordinator_hidden), | |
| nn.GELU(), | |
| nn.LayerNorm(config.coordinator_hidden), | |
| ) | |
| self.gate_head = nn.Linear(config.coordinator_hidden, 4) # think, code, mem, scratch | |
| self.step_head = nn.Linear(config.coordinator_hidden, config.max_steps + 1) | |
| def forward(self, hidden: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| pooled = hidden.mean(1) # (B, H) | |
| feat = self.hidden(pooled) | |
| gates = torch.sigmoid(self.gate_head(feat)) # (B, 4) | |
| steps_dist = torch.softmax(self.step_head(feat), dim=-1) # (B, max_steps+1) | |
| steps = torch.argmax(steps_dist, dim=-1) # (B,) | |
| return gates, steps_dist, steps | |
| # --------------------------------------------------------------------------- | |
| # System-2: iterative reasoning loop | |
| # --------------------------------------------------------------------------- | |
| class _ReasonerLayer(nn.Module): | |
| """Single self-attention + FFN layer, weight-tied across loop iterations.""" | |
| def __init__(self, dim: int, heads: int, ff: int): | |
| super().__init__() | |
| self.dim = dim | |
| self.heads = heads | |
| self.head_dim = dim // heads | |
| self.norm1 = nn.LayerNorm(dim) | |
| self.qkv = nn.Linear(dim, 3 * dim) | |
| self.out_proj = nn.Linear(dim, dim) | |
| self.norm2 = nn.LayerNorm(dim) | |
| self.ff = nn.Sequential(nn.Linear(dim, ff), nn.GELU(), nn.Linear(ff, dim)) | |
| self.ff_norm = nn.LayerNorm(dim) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| B, T, D = x.shape | |
| h, hd = self.heads, self.head_dim | |
| qkv = self.qkv(self.norm1(x)).reshape(B, T, 3, h, hd).permute(2, 0, 3, 1, 4) | |
| q, k, v = qkv[0], qkv[1], qkv[2] | |
| attn = torch.matmul(q, k.transpose(-1, -2)) / math.sqrt(hd) | |
| attn = F.softmax(attn, dim=-1) | |
| out = torch.matmul(attn, v).transpose(1, 2).reshape(B, T, D) | |
| x = x + self.out_proj(out) | |
| x = x + self.ff(self.ff_norm(self.norm2(x))) | |
| return x | |
| class MultiStepReasoner(nn.Module): | |
| """ | |
| System-2 thinking loop. Compresses hidden states to a small workspace, | |
| refines them through a weight-tied attention layer `steps` times, then | |
| projects back. Produces a scratchpad of intermediate states. | |
| """ | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| dim = config.reasoner_dim | |
| self.in_proj = nn.Linear(hidden_dim, dim) | |
| self.out_proj = nn.Linear(dim, hidden_dim) | |
| self.layer = _ReasonerLayer(dim, config.reasoner_heads, config.reasoner_ff) | |
| self.step_emb = nn.Embedding(config.max_steps + 1, dim) | |
| nn.init.normal_(self.step_emb.weight, std=0.02) | |
| nn.init.zeros_(self.out_proj.weight) | |
| nn.init.zeros_(self.out_proj.bias) | |
| self.max_steps = config.max_steps | |
| def forward(self, hidden: torch.Tensor, steps: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: | |
| x = self.in_proj(hidden) # (B, T, dim) | |
| x = x + self.step_emb(torch.zeros_like(steps).long()).unsqueeze(1) # step-0 token | |
| scratch = [x] | |
| for s in range(1, self.max_steps + 1): | |
| active = steps >= s # (B,) which rows still think | |
| if active.any(): | |
| xa = x + self.step_emb(torch.full_like(steps, s).long()).unsqueeze(1) | |
| x = torch.where(active.unsqueeze(1).unsqueeze(1), self.layer(xa), x) | |
| scratch.append(x) | |
| else: | |
| scratch.append(x) | |
| out = self.out_proj(x) # (B, T, H) | |
| return out, scratch[-1] | |
| class MultiHeadCoT(nn.Module): | |
| """Multi-head chain-of-thought reasoning: generates N parallel reasoning paths | |
| and fuses them for higher accuracy on complex tasks.""" | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| self.num_heads = config.num_cot_heads | |
| cot_dim = config.cot_hidden | |
| self.heads = nn.ModuleList([ | |
| nn.Sequential( | |
| nn.Linear(hidden_dim, cot_dim), | |
| nn.GELU(), | |
| nn.LayerNorm(cot_dim), | |
| nn.Linear(cot_dim, hidden_dim), | |
| ) for _ in range(self.num_heads) | |
| ]) | |
| self.fusion = nn.Sequential( | |
| nn.Linear(hidden_dim * (self.num_heads + 1), hidden_dim), | |
| nn.GELU(), | |
| nn.LayerNorm(hidden_dim), | |
| nn.Linear(hidden_dim, hidden_dim), | |
| ) | |
| nn.init.zeros_(self.fusion[-1].weight) | |
| nn.init.zeros_(self.fusion[-1].bias) | |
| def forward(self, hidden: torch.Tensor) -> torch.Tensor: | |
| B, T, H = hidden.shape | |
| paths = [hidden] | |
| for head in self.heads: | |
| paths.append(head(hidden)) | |
| fused = self.fusion(torch.cat(paths, dim=-1)) | |
| return hidden + fused # Residual connection | |
| # --------------------------------------------------------------------------- | |
| # Code structure awareness | |
| # --------------------------------------------------------------------------- | |
| class CodeAwareBias(nn.Module): | |
| """ | |
| Injects code structure as a learned bias. `code_feat` holds per-token | |
| features (is_code, indent depth, bracket balance, newline). A small net | |
| maps them to a per-token weight that scales a projected hidden state, | |
| so the model can pay structural attention to indentation and brackets. | |
| """ | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| self.structure_net = nn.Sequential( | |
| nn.Linear(config.code_feat_dim, 32), | |
| nn.GELU(), | |
| nn.Linear(32, 1), | |
| ) | |
| self.proj = nn.Linear(hidden_dim, hidden_dim) | |
| nn.init.zeros_(self.proj.weight) | |
| nn.init.zeros_(self.proj.bias) | |
| def forward(self, hidden: torch.Tensor, code_feat: Optional[torch.Tensor]) -> torch.Tensor: | |
| if code_feat is None: | |
| return hidden | |
| code_feat = code_feat.to(hidden.dtype) | |
| w = torch.sigmoid(self.structure_net(code_feat)) # (B, T, 1) | |
| return hidden + self.proj(hidden) * w | |
| def build_code_features(tokenizer, input_ids: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Build per-token code-structure features (B, T, 4) from token strings: | |
| [0] is_code_like (indent / brackets / operators / newlines) | |
| [1] indent_depth (normalized leading whitespace) | |
| [2] bracket_balance (+1 open, 0 neutral, -1 close -> mapped to 0/0.5/1) | |
| [3] has_newline | |
| """ | |
| code_chars = set("{}[]();=<>!&|+-*/%'\"`#@.,:") | |
| feats = [] | |
| for row in input_ids.tolist(): | |
| tokens = tokenizer.convert_ids_to_tokens(row) | |
| row_feats = [] | |
| for tok in tokens: | |
| is_code = any(c in code_chars for c in tok) | |
| indent = 0.0 | |
| stripped = tok.lstrip() | |
| if stripped and tok != stripped: | |
| indent = min((len(tok) - len(stripped)) / 8.0, 1.0) | |
| is_code = True | |
| bal = 0.0 | |
| if any(c in "{[(" for c in tok): | |
| bal = 1.0 | |
| elif any(c in "}])" for c in tok): | |
| bal = 0.0 | |
| else: | |
| bal = 0.5 | |
| newline = 1.0 if "\n" in tok else 0.0 | |
| row_feats.append([1.0 if is_code else 0.0, indent, bal, newline]) | |
| # pad/truncate to input length | |
| feats.append(row_feats[: input_ids.shape[1]]) | |
| # pad rows to same length | |
| max_len = max(len(r) for r in feats) | |
| padded = [ | |
| r + [[0.0, 0.0, 0.5, 0.0]] * (max_len - len(r)) | |
| for r in feats | |
| ] | |
| return torch.tensor(padded, dtype=torch.float32) | |
| # --------------------------------------------------------------------------- | |
| # Scratchpad (cross-turn working memory) + Verifier | |
| # --------------------------------------------------------------------------- | |
| class ScratchpadMemory(nn.Module): | |
| """ | |
| Cross-turn working memory in the full hidden-dim space. Writes the last | |
| reasoning state and reads it back on the next call, so long reasoning can | |
| continue across assistant turns. | |
| """ | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| self.key_proj = nn.Linear(hidden_dim, hidden_dim) | |
| self.read_proj = nn.Linear(hidden_dim, hidden_dim) | |
| self.state = None | |
| nn.init.zeros_(self.read_proj.weight) | |
| nn.init.zeros_(self.read_proj.bias) | |
| def read(self, hidden: torch.Tensor) -> torch.Tensor: | |
| """Returns a bias added to the current refined hidden state.""" | |
| if self.state is None: | |
| return torch.zeros_like(hidden) | |
| bias = self.read_proj(self.state) # (H,) | |
| return bias.unsqueeze(0).unsqueeze(0) # (1, 1, H) | |
| def write(self, hidden: torch.Tensor): | |
| with torch.no_grad(): | |
| self.state = self.key_proj(hidden.detach().mean(1)).mean(0) # (H,) | |
| class VerifierHead(nn.Module): | |
| """ | |
| Lightweight self-critique scorer. Scores a full sequence with a scalar; | |
| trained to match normalized sequence likelihood. Used for best-of-n | |
| decoding: generate several candidates, keep the highest-scoring one. | |
| """ | |
| def __init__(self, hidden_dim: int): | |
| super().__init__() | |
| self.net = nn.Sequential( | |
| nn.Linear(hidden_dim, 128), | |
| nn.GELU(), | |
| nn.Linear(128, 1), | |
| ) | |
| def forward(self, hidden: torch.Tensor) -> torch.Tensor: | |
| pooled = hidden.mean(1) # (B, H) | |
| return self.net(pooled).squeeze(-1) # (B,) | |
| class QuantizedMemoryModule(MemoryModule): | |
| """MemoryModule with INT8/INT4 quantized KV cache for memory efficiency.""" | |
| def __init__(self, memory_size: int, memory_dim: int, hidden_dim: int, | |
| quantize: bool = True, quant_bits: int = 8): | |
| super().__init__(memory_size, memory_dim, hidden_dim) | |
| self.quantize = quantize | |
| self.quant_bits = quant_bits | |
| self._quant_scale = None | |
| def _quantize(self, x: torch.Tensor) -> torch.Tensor: | |
| if not self.quantize or self.quant_bits >= 16: | |
| return x | |
| scale = x.abs().max() / (2 ** (self.quant_bits - 1) - 1) | |
| self._quant_scale = scale.item() | |
| q = torch.round(x / scale).clamp(-(2 ** (self.quant_bits - 1)), 2 ** (self.quant_bits - 1) - 1) | |
| return (q * scale).to(x.dtype) | |
| def read(self, hidden: torch.Tensor) -> torch.Tensor: | |
| if self.quantize and self.mem_k_buf is not None: | |
| self.mem_k_buf = self._quantize(self.mem_k_buf) | |
| self.mem_v_buf = self._quantize(self.mem_v_buf) | |
| return super().read(hidden) | |
| def write(self, hidden: torch.Tensor): | |
| super().write(hidden) | |
| class MixtureOfDepths(nn.Module): | |
| """MoD: per-token gating to dynamically skip transformer layers. | |
| Uses a lightweight router with temperature annealing for adaptive layer skipping, | |
| reducing compute by ~30-50% with minimal accuracy loss. | |
| """ | |
| def __init__(self, hidden_dim: int, mod_hidden: int, keep_prob: float = 0.8, | |
| dropout: float = 0.1, temperature: float = 1.0, temperature_anneal: float = 0.995): | |
| super().__init__() | |
| self.keep_prob = keep_prob | |
| self.temperature = temperature | |
| self.temperature_anneal = temperature_anneal | |
| self.router = nn.Sequential( | |
| nn.Linear(hidden_dim, mod_hidden), | |
| nn.GELU(), | |
| nn.LayerNorm(mod_hidden), | |
| nn.Linear(mod_hidden, 1), | |
| ) | |
| self.dropout = nn.Dropout(dropout) | |
| def forward(self, hidden: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: | |
| B, T, H = hidden.shape | |
| logits = self.router(hidden.detach()) # (B, T, 1) | |
| probs = torch.sigmoid(logits / self.temperature) | |
| mask = torch.bernoulli(probs * 0.5 + self.keep_prob * 0.5).expand_as(hidden) | |
| mask = self.dropout(mask) | |
| if self.training: | |
| self.temperature = max(0.1, self.temperature * self.temperature_anneal) | |
| return hidden * mask, probs | |
| class MemoryEfficientAttention(nn.Module): | |
| """Memory-efficient attention using PyTorch 2.0+ SDPA with optional Flash Attention. | |
| Falls back to standard attention if SDPA is unavailable. | |
| """ | |
| def __init__(self, dim: int, heads: int, dropout: float = 0.0): | |
| super().__init__() | |
| self.dim = dim | |
| self.heads = heads | |
| self.head_dim = dim // heads | |
| self.scale = self.head_dim ** -0.5 | |
| self.qkv = nn.Linear(dim, 3 * dim, bias=False) | |
| self.out_proj = nn.Linear(dim, dim, bias=False) | |
| self.dropout_p = dropout | |
| self.use_sdpa = hasattr(F, 'scaled_dot_product_attention') | |
| def forward(self, x: torch.Tensor, attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor: | |
| B, T, D = x.shape | |
| h, hd = self.heads, self.head_dim | |
| qkv = self.qkv(x).reshape(B, T, 3, h, hd).permute(2, 0, 3, 1, 4) | |
| q, k, v = qkv[0], qkv[1], qkv[2] | |
| if self.use_sdpa: | |
| try: | |
| out = F.scaled_dot_product_attention( | |
| q, k, v, | |
| attn_mask=attention_mask, | |
| dropout_p=self.dropout_p if self.training else 0.0, | |
| is_causal=(attention_mask is None), | |
| ) | |
| out = out.transpose(1, 2).reshape(B, T, D) | |
| return self.out_proj(out) | |
| except Exception: | |
| pass | |
| attn = torch.matmul(q, k.transpose(-1, -2)) * self.scale | |
| if attention_mask is not None: | |
| attn = attn + attention_mask | |
| attn = F.softmax(attn, dim=-1) | |
| if self.training and self.dropout_p > 0: | |
| attn = F.dropout(attn, p=self.dropout_p) | |
| out = torch.matmul(attn, v).transpose(1, 2).reshape(B, T, D) | |
| return self.out_proj(out) | |
| class DynamicMoEBlock(nn.Module): | |
| """Sparse MoE with dynamic expert expansion, pruning, and load-balancing.""" | |
| def __init__(self, hidden_dim: int, num_experts: int, expert_hidden: int, | |
| top_k: int, prune_threshold: float = 0.02, expand_threshold: float = 0.15, max_experts: int = 64): | |
| super().__init__() | |
| self.num_experts = num_experts | |
| self.max_experts = max_experts | |
| self.top_k = top_k | |
| self.prune_threshold = prune_threshold | |
| self.expand_threshold = expand_threshold | |
| self.gate = nn.Linear(hidden_dim, num_experts, bias=False) | |
| self.experts = nn.ModuleList([ | |
| Expert(hidden_dim, expert_hidden) for _ in range(num_experts) | |
| ]) | |
| self.expert_usage = torch.zeros(num_experts) | |
| self._pruned = set() | |
| self._expansion_count = 0 | |
| def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: | |
| B, T, H = x.shape | |
| flat = x.reshape(-1, H) | |
| gate_logits = self.gate(flat) | |
| probs = F.softmax(gate_logits, dim=-1) | |
| topk_vals, topk_idx = torch.topk(gate_logits, self.top_k, dim=-1) | |
| topk_vals = F.softmax(topk_vals, dim=-1) | |
| routing = torch.zeros_like(probs) | |
| routing.scatter_(1, topk_idx, topk_vals) | |
| out = torch.zeros_like(flat) | |
| for i, expert in enumerate(self.experts): | |
| if i in self._pruned: | |
| continue | |
| sel = routing[:, i] > 0 | |
| if sel.any(): | |
| out[sel] += routing[sel, i].unsqueeze(-1) * expert(flat[sel]) | |
| if i < len(self.expert_usage): | |
| self.expert_usage[i] += sel.sum().item() | |
| f_i = routing.mean(0) | |
| P_i = probs.mean(0) | |
| aux = (f_i * P_i).sum() * self.num_experts | |
| return out.view(B, T, H), aux | |
| def prune_and_expand_experts(self): | |
| """Dynamically prune underused experts and clone overused ones.""" | |
| total = self.expert_usage.sum() | |
| if total == 0: | |
| self.expert_usage.zero_() | |
| return | |
| usage_ratios = self.expert_usage / total | |
| active_experts = [i for i in range(len(self.experts)) if i not in self._pruned] | |
| for i in active_experts: | |
| if usage_ratios[i] < self.prune_threshold and len(self._pruned) < len(self.experts) - 1: | |
| self._pruned.add(i) | |
| print(f"Pruned expert {i} (usage {usage_ratios[i]:.4f})") | |
| if len(self.experts) < self.max_experts: | |
| avg_usage = usage_ratios[active_experts].mean().item() | |
| for i in active_experts: | |
| if usage_ratios[i] > self.expand_threshold and len(self.experts) < self.max_experts: | |
| new_expert = Expert( | |
| self.experts[i].in_proj.in_features, | |
| self.experts[i].in_proj.out_features | |
| ) | |
| new_expert.load_state_dict(self.experts[i].state_dict()) | |
| with torch.no_grad(): | |
| for param in new_expert.parameters(): | |
| param.add_(torch.randn_like(param) * 0.01) | |
| self.experts.append(new_expert) | |
| self.expert_usage = torch.cat([self.expert_usage, torch.zeros(1)]) | |
| self._expansion_count += 1 | |
| print(f"Expanded expert {i} -> new expert {len(self.experts)-1}") | |
| self.expert_usage.zero_() | |
| print(f"Active experts: {len(self.experts) - len(self._pruned)}/{len(self.experts)}") | |
| class MultimodalFusion(nn.Module): | |
| """Fuse text + vision + audio + video embeddings into a unified representation.""" | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| self.vision_proj = nn.Linear(config.vision_dim, hidden_dim) | |
| self.audio_proj = nn.Linear(config.audio_dim, hidden_dim) | |
| self.video_proj = nn.Linear(config.video_dim, hidden_dim) | |
| self.fusion = nn.Sequential( | |
| nn.Linear(hidden_dim * 4, config.fusion_hidden), | |
| nn.GELU(), | |
| nn.LayerNorm(config.fusion_hidden), | |
| nn.Linear(config.fusion_hidden, hidden_dim), | |
| ) | |
| nn.init.zeros_(self.fusion[-1].weight) | |
| nn.init.zeros_(self.fusion[-1].bias) | |
| def forward(self, text: torch.Tensor, vision: Optional[torch.Tensor] = None, | |
| audio: Optional[torch.Tensor] = None, video: Optional[torch.Tensor] = None) -> torch.Tensor: | |
| parts = [text] | |
| if vision is not None: | |
| parts.append(self.vision_proj(vision)) | |
| if audio is not None: | |
| parts.append(self.audio_proj(audio)) | |
| if video is not None: | |
| parts.append(self.video_proj(video)) | |
| while len(parts) < 4: | |
| parts.append(torch.zeros_like(text)) | |
| fused = self.fusion(torch.cat(parts, dim=-1)) | |
| return text + fused | |
| class ToolUseModule(nn.Module): | |
| """JSON-structured function calling with validation and execution.""" | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| self.max_tools = config.max_tools | |
| self.tool_embeddings = nn.Embedding(config.max_tools, hidden_dim) | |
| self.tool_classifier = nn.Sequential( | |
| nn.Linear(hidden_dim, config.tool_hidden), | |
| nn.GELU(), | |
| nn.Linear(config.tool_hidden, config.max_tools), | |
| ) | |
| self.arg_proj = nn.Linear(hidden_dim, hidden_dim) | |
| nn.init.normal_(self.tool_embeddings.weight, std=0.02) | |
| def forward(self, hidden: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: | |
| pooled = hidden.mean(1) | |
| tool_logits = self.tool_classifier(pooled) | |
| tool_probs = F.softmax(tool_logits, dim=-1) | |
| tool_idx = torch.argmax(tool_probs, dim=-1) | |
| tool_emb = self.tool_embeddings(tool_idx) | |
| args = self.arg_proj(pooled) | |
| return tool_emb, args | |
| def generate_tool_call(self, hidden: torch.Tensor, tokenizer) -> str: | |
| """Generate a JSON tool call from hidden state.""" | |
| tool_emb, args = self.forward(hidden) | |
| tool_idx = torch.argmax(self.tool_classifier(hidden.mean(1)), dim=-1).item() | |
| tool_name = f"tool_{tool_idx}" | |
| arg_vec = args[0].detach().cpu().numpy().tolist() | |
| return json.dumps({ | |
| "tool": tool_name, | |
| "arguments": {"vector": arg_vec[:10]}, | |
| "confidence": float(torch.softmax(self.tool_classifier(hidden.mean(1)), dim=-1)[0, tool_idx].item()) | |
| }) | |
| class DocumentModule(nn.Module): | |
| """PDF/DOCX/OCR with layout-aware parsing for document understanding.""" | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| self.max_pages = config.doc_max_pages | |
| self.page_proj = nn.Linear(hidden_dim, config.doc_hidden) | |
| self.layout_encoder = nn.Sequential( | |
| nn.Linear(config.doc_hidden + 4, config.doc_hidden), | |
| nn.GELU(), | |
| nn.Linear(config.doc_hidden, hidden_dim), | |
| ) | |
| self.out_proj = nn.Linear(config.doc_hidden, hidden_dim) | |
| nn.init.zeros_(self.out_proj.weight) | |
| nn.init.zeros_(self.out_proj.bias) | |
| nn.init.zeros_(self.layout_encoder[-1].weight) | |
| nn.init.zeros_(self.layout_encoder[-1].bias) | |
| def forward(self, hidden: torch.Tensor, layout_info: Optional[torch.Tensor] = None) -> torch.Tensor: | |
| B, T, H = hidden.shape | |
| page_emb = self.page_proj(hidden) | |
| if layout_info is not None: | |
| layout = layout_info.to(hidden.dtype) | |
| page_emb = self.layout_encoder(torch.cat([page_emb, layout], dim=-1)) | |
| else: | |
| page_emb = self.out_proj(page_emb) | |
| return hidden + page_emb | |
| def extract_text(self, source) -> str: | |
| """Extract text from PDF/DOCX/image with OCR fallback.""" | |
| try: | |
| if hasattr(source, 'endswith') and source.endswith('.pdf'): | |
| return self._extract_pdf(source) | |
| elif hasattr(source, 'endswith') and source.endswith('.docx'): | |
| return self._extract_docx(source) | |
| else: | |
| return self._extract_image_ocr(source) | |
| except Exception as e: | |
| return f"[document extraction error: {e}]" | |
| def _extract_pdf(self, path: str) -> str: | |
| try: | |
| import fitz | |
| doc = fitz.open(path) | |
| pages = [] | |
| for i in range(min(len(doc), self.max_pages)): | |
| pages.append(doc[i].get_text()) | |
| return "\n\n".join(pages) | |
| except ImportError: | |
| return "[PDF extraction requires PyMuPDF: pip install pymupdf]" | |
| def _extract_docx(self, path: str) -> str: | |
| try: | |
| import docx2txt | |
| return docx2txt.process(path) | |
| except ImportError: | |
| return "[DOCX extraction requires docx2txt: pip install docx2txt]" | |
| def _extract_image_ocr(self, source) -> str: | |
| try: | |
| import pytesseract | |
| from PIL import Image | |
| img = Image.open(source) | |
| return pytesseract.image_to_string(img) | |
| except ImportError: | |
| return "[OCR requires pytesseract + Pillow: pip install pytesseract pillow]" | |
| class VideoModule(nn.Module): | |
| """Temporal frame sampling + motion features for video understanding.""" | |
| def __init__(self, config: MorphConfig, hidden_dim: int): | |
| super().__init__() | |
| self.max_frames = config.video_max_frames | |
| self.frame_proj = nn.Linear(hidden_dim, config.video_hidden) | |
| self.temporal_encoder = nn.GRU( | |
| config.video_hidden, config.video_hidden, | |
| batch_first=True, bidirectional=False | |
| ) | |
| self.motion_proj = nn.Linear(config.video_hidden, hidden_dim) | |
| nn.init.zeros_(self.motion_proj.weight) | |
| nn.init.zeros_(self.motion_proj.bias) | |
| def forward(self, hidden: torch.Tensor, frame_embeddings: Optional[torch.Tensor] = None) -> torch.Tensor: | |
| B, T, H = hidden.shape | |
| if frame_embeddings is None: | |
| return hidden | |
| frame_emb = self.frame_proj(frame_embeddings) | |
| _, last_hidden = self.temporal_encoder(frame_emb) | |
| motion = self.motion_proj(last_hidden.squeeze(0)) | |
| return hidden + motion.unsqueeze(1) | |
| class CodeSandbox: | |
| """Safe Python code execution with AST validation and resource limits.""" | |
| def __init__(self, timeout: float = 5.0, max_memory_mb: int = 128): | |
| self.timeout = timeout | |
| self.max_memory = max_memory_mb | |
| self._allowed_modules = { | |
| 'math', 'random', 'datetime', 'collections', 'itertools', | |
| 'functools', 'operator', 'statistics', 'json', 're', | |
| 'string', 'typing', 'copy', 'heapq', 'bisect', 'array', | |
| } | |
| self._allowed_builtins = { | |
| 'print', 'len', 'range', 'enumerate', 'zip', 'map', 'filter', | |
| 'sum', 'min', 'max', 'abs', 'round', 'sorted', 'list', 'dict', | |
| 'set', 'tuple', 'int', 'float', 'str', 'bool', 'bytes', | |
| 'True', 'False', 'None', 'isinstance', 'type', 'hasattr', | |
| 'getattr', 'setattr', 'property', 'staticmethod', 'classmethod', | |
| } | |
| def validate_ast(self, code: str) -> Tuple[bool, str]: | |
| """Check code for unsafe operations using AST analysis.""" | |
| import ast | |
| try: | |
| tree = ast.parse(code) | |
| except SyntaxError as e: | |
| return False, f"Syntax error: {e}" | |
| for node in ast.walk(tree): | |
| if isinstance(node, ast.Import): | |
| for alias in node.names: | |
| if alias.name.split('.')[0] not in self._allowed_modules: | |
| return False, f"Import of '{alias.name}' not allowed" | |
| elif isinstance(node, ast.ImportFrom): | |
| if node.module and node.module.split('.')[0] not in self._allowed_modules: | |
| return False, f"Import from '{node.module}' not allowed" | |
| elif hasattr(ast, 'Exec') and isinstance(node, ast.Exec): | |
| return False, "exec() is not allowed" | |
| elif hasattr(ast, 'Eval') and isinstance(node, ast.Eval): | |
| return False, "eval() is not allowed" | |
| elif isinstance(node, ast.Call): | |
| func = node.func | |
| if isinstance(func, ast.Name) and func.id in ('eval', 'exec', '__import__', 'open', 'compile'): | |
| return False, f"'{func.id}()' is not allowed" | |
| return True, "OK" | |
| def execute(self, code: str, context: Optional[dict] = None) -> dict: | |
| """Execute code in a restricted environment.""" | |
| import traceback | |
| safe, msg = self.validate_ast(code) | |
| if not safe: | |
| return {"success": False, "output": "", "error": msg} | |
| safe_globals = {"__builtins__": {k: __builtins__[k] for k in self._allowed_builtins if k in __builtins__}} | |
| safe_locals = context or {} | |
| try: | |
| result = eval(code, safe_globals, safe_locals) | |
| return {"success": True, "output": str(result), "error": ""} | |
| except Exception as e: | |
| return {"success": False, "output": "", "error": traceback.format_exc()} | |
| # --------------------------------------------------------------------------- | |
| # MorphModel | |
| # --------------------------------------------------------------------------- | |
| class MorphModel(nn.Module): | |
| """ | |
| MORPH-AI v6. Base model + 15 novel subsystems wired into the logits. | |
| Forward path: | |
| embeds = base.embed_tokens(input_ids) [+ skill injection] | |
| base_hidden = base.layers(embeds) # frozen + LoRA + memory-efficient attention | |
| mod_mask, mod_probs = MixtureOfDepths(base_hidden) # dynamic layer skip | |
| hidden = base_hidden * mod_mask # MoD gated | |
| gates, steps_dist, steps = Coordinator(hidden) | |
| reasoned, scratch = MultiStepReasoner(hidden, steps) # System 2 | |
| code_bias = CodeAwareBias(reasoned, code_feat) # if code gate | |
| depth_emb, depth_dist = DepthEmbeddings(reasoned) | |
| refined = reasoned + depth_emb.unsqueeze(1) | |
| fused = MultimodalFusion(refined, vision, audio, video) # multimodal | |
| moe_out, moe_aux = DynamicMoEBlock(fused) # sparse MoE + pruning | |
| mem_out = QuantizedMemory.read(fused) # quantized KV memory | |
| scratch_out = Scratchpad.read(scratch) # cross-turn memory | |
| tool_emb, args = ToolUseModule(fused) # tool calling | |
| doc_out = DocumentModule(fused, layout_info) # document understanding | |
| video_out = VideoModule(fused, frame_embeddings) # video understanding | |
| final = fused + think*(moe_out+mem_out) + code_bias + scratch_out + doc_out + video_out | |
| logits = base.lm_head(final) | |
| score = VerifierHead(final) # self-critique | |
| """ | |
| def __init__(self, config: Optional[MorphConfig] = None): | |
| super().__init__() | |
| self.cfg = config or MorphConfig() | |
| self.novel_trained = False | |
| print(f"Loading base model: {self.cfg.base_model}") | |
| try: | |
| self.base_model_raw = AutoModelForCausalLM.from_pretrained( | |
| self.cfg.base_model, | |
| dtype=torch.bfloat16, | |
| device_map="auto", | |
| trust_remote_code=True, | |
| ) | |
| except TypeError: | |
| self.base_model_raw = AutoModelForCausalLM.from_pretrained( | |
| self.cfg.base_model, | |
| torch_dtype=torch.bfloat16, | |
| device_map="auto", | |
| trust_remote_code=True, | |
| ) | |
| self.tokenizer = AutoTokenizer.from_pretrained(self.cfg.base_model, trust_remote_code=True) | |
| if self.tokenizer.pad_token is None: | |
| self.tokenizer.pad_token = self.tokenizer.eos_token | |
| hidden_dim = self.base_model_raw.config.hidden_size | |
| vocab_size = self.base_model_raw.config.vocab_size | |
| # Extend context window via RoPE scaling if configured | |
| original_max = getattr(self.base_model_raw.config, 'max_position_embeddings', 2048) | |
| if self.cfg.max_seq_len > original_max: | |
| print(f"Extending context: {original_max} -> {self.cfg.max_seq_len}") | |
| if hasattr(self.base_model_raw.config, 'rope_scaling') and self.base_model_raw.config.rope_scaling is None: | |
| self.base_model_raw.config.rope_scaling = { | |
| "type": "yarn", | |
| "factor": self.cfg.max_seq_len / original_max, | |
| } | |
| self.base_model_raw.config.max_position_embeddings = self.cfg.max_seq_len | |
| self.tokenizer.model_max_length = self.cfg.max_seq_len | |
| # v6 subsystems | |
| self.coordinator = Coordinator(self.cfg, hidden_dim) | |
| self.reasoner = MultiStepReasoner(self.cfg, hidden_dim) | |
| self.code_bias = CodeAwareBias(self.cfg, hidden_dim) | |
| self.scratchpad = ScratchpadMemory(self.cfg, hidden_dim) | |
| self.verifier = VerifierHead(hidden_dim) | |
| self.skill_module = SkillTokenModule(self.cfg, hidden_dim) | |
| self.depth_module = DepthEmbeddings(self.cfg, hidden_dim) | |
| self.moe_block = DynamicMoEBlock( | |
| hidden_dim, self.cfg.num_experts, self.cfg.expert_hidden, | |
| self.cfg.moe_top_k, self.cfg.moe_prune_threshold, | |
| self.cfg.moe_expand_threshold, self.cfg.max_experts | |
| ) | |
| self.memory = QuantizedMemoryModule(self.cfg.memory_size, self.cfg.memory_dim, hidden_dim, self.cfg.memory_quantize, self.cfg.memory_quant_bits) | |
| self.mod = MixtureOfDepths( | |
| hidden_dim, self.cfg.mod_hidden, self.cfg.mod_keep_prob, | |
| self.cfg.mod_dropout, self.cfg.mod_temperature, self.cfg.mod_temperature_anneal | |
| ) | |
| self.multimodal_fusion = MultimodalFusion(self.cfg, hidden_dim) | |
| self.tool_use = ToolUseModule(self.cfg, hidden_dim) | |
| self.document_module = DocumentModule(self.cfg, hidden_dim) | |
| self.video_module = VideoModule(self.cfg, hidden_dim) | |
| self.code_sandbox = CodeSandbox(self.cfg.sandbox_timeout, self.cfg.sandbox_max_memory) | |
| self.cot_reasoner = MultiHeadCoT(self.cfg, hidden_dim) | |
| # cast novel components to the base model's compute dtype | |
| self._dtype = self.base_model_raw.model.embed_tokens.weight.dtype | |
| for mod in ( | |
| self.coordinator, self.reasoner, self.code_bias, self.scratchpad, | |
| self.verifier, self.skill_module, self.depth_module, self.moe_block, | |
| self.memory, self.mod, self.multimodal_fusion, self.tool_use, | |
| self.document_module, self.video_module, self.code_sandbox, self.cot_reasoner, | |
| ): | |
| mod.to(self._dtype) | |
| self.vocab_size = vocab_size | |
| self.base_model = None | |
| self._skill_lora_modules: Dict[str, nn.Module] = {} | |
| self._plugins: Dict[str, nn.Module] = {} | |
| # Load plugins from plugin_dir if specified | |
| if self.cfg.plugin_dir: | |
| self.load_plugins(self.cfg.plugin_dir) | |
| def load_plugins(self, plugin_dir: str): | |
| """Load custom capability plugins from a directory.""" | |
| import os | |
| import importlib.util | |
| plugin_path = Path(plugin_dir) | |
| if not plugin_path.exists(): | |
| print(f"Plugin directory not found: {plugin_dir}") | |
| return | |
| for file in plugin_path.glob("*.py"): | |
| if file.name.startswith("_"): | |
| continue | |
| try: | |
| spec = importlib.util.spec_from_file_location(file.stem, file) | |
| mod = importlib.util.module_from_spec(spec) | |
| spec.loader.exec_module(mod) | |
| for attr_name in dir(mod): | |
| attr = getattr(mod, attr_name) | |
| if isinstance(attr, type) and issubclass(attr, nn.Module) and attr is not nn.Module: | |
| plugin_name = getattr(attr, 'plugin_name', attr_name) | |
| plugin_instance = attr(self.cfg, hidden_dim=self.base_model_raw.config.hidden_size) | |
| setattr(self, f"plugin_{plugin_name}", plugin_instance) | |
| self._plugins[plugin_name] = plugin_instance | |
| plugin_instance.to(self._dtype) | |
| print(f"Loaded plugin: {plugin_name} from {file.name}") | |
| except Exception as e: | |
| print(f"Failed to load plugin {file.name}: {e}") | |
| self._plugins: Dict[str, nn.Module] = {} | |
| # Load plugins from plugin_dir if specified | |
| if self.cfg.plugin_dir: | |
| self.load_plugins(self.cfg.plugin_dir) | |
| # ---- gradient-checkpointing passthrough (Trainer calls these on the top model) ---- | |
| def gradient_checkpointing_enable(self, gradient_checkpointing_kwargs=None): | |
| target = self.base_model or self.base_model_raw | |
| if hasattr(target, "gradient_checkpointing_enable"): | |
| return target.gradient_checkpointing_enable( | |
| gradient_checkpointing_kwargs=gradient_checkpointing_kwargs | |
| ) | |
| def gradient_checkpointing_disable(self): | |
| target = self.base_model or self.base_model_raw | |
| if hasattr(target, "gradient_checkpointing_disable"): | |
| return target.gradient_checkpointing_disable() | |
| def enable_input_require_grads(self): | |
| target = self.base_model or self.base_model_raw | |
| if hasattr(target, "enable_input_require_grads"): | |
| return target.enable_input_require_grads() | |
| def disable_input_require_grads(self): | |
| target = self.base_model or self.base_model_raw | |
| if hasattr(target, "disable_input_require_grads"): | |
| return target.disable_input_require_grads() | |
| # ---- LoRA / PEFT ---- | |
| def apply_lora(self, target_modules: Optional[List[str]] = None): | |
| target_modules = target_modules or [ | |
| "q_proj", "k_proj", "v_proj", "o_proj", | |
| "gate_proj", "up_proj", "down_proj", | |
| ] | |
| lora_config = LoraConfig( | |
| r=self.cfg.lora_rank, | |
| lora_alpha=self.cfg.lora_alpha, | |
| lora_dropout=self.cfg.lora_dropout, | |
| target_modules=target_modules, | |
| task_type=TaskType.CAUSAL_LM, | |
| bias="none", | |
| ) | |
| self.base_model = get_peft_model(self.base_model_raw, lora_config) | |
| self.base_model.print_trainable_parameters() | |
| return self.base_model | |
| # ---- forward ---- | |
| def forward( | |
| self, | |
| input_ids: Optional[torch.Tensor] = None, | |
| inputs_embeds: Optional[torch.Tensor] = None, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| labels: Optional[torch.Tensor] = None, | |
| skill_indices: Optional[torch.Tensor] = None, | |
| code_feat: Optional[torch.Tensor] = None, | |
| force_depth: Optional[int] = None, | |
| use_adaptive: bool = True, | |
| vision_embeds: Optional[torch.Tensor] = None, | |
| audio_embeds: Optional[torch.Tensor] = None, | |
| video_embeds: Optional[torch.Tensor] = None, | |
| frame_embeddings: Optional[torch.Tensor] = None, | |
| layout_info: Optional[torch.Tensor] = None, | |
| **kwargs, | |
| ): | |
| if self.base_model is None: | |
| raise RuntimeError("Call apply_lora() before forward().") | |
| if inputs_embeds is None: | |
| inputs_embeds = self.base_model_raw.model.embed_tokens(input_ids) | |
| skill_emb = self.skill_module(skill_indices) | |
| if skill_emb is not None: | |
| inputs_embeds = inputs_embeds + 0.1 * skill_emb.unsqueeze(1) | |
| base_out = self.base_model( | |
| inputs_embeds=inputs_embeds, | |
| attention_mask=attention_mask, | |
| output_hidden_states=True, | |
| ) | |
| base_hidden = base_out.hidden_states[-1] # (B, T, H) | |
| # If the novel v4 components were never trained, skip perturbations. | |
| if not getattr(self, "novel_trained", True): | |
| logits = self.base_model_raw.lm_head(base_hidden.to(self.base_model_raw.lm_head.weight.dtype)) | |
| out = ModelOutput( | |
| logits=logits, | |
| hidden_states=base_hidden, | |
| refined=base_hidden, | |
| gates=torch.zeros((base_hidden.shape[0], 4), device=base_hidden.device), | |
| steps_dist=torch.zeros((base_hidden.shape[0], self.cfg.max_steps), device=base_hidden.device), | |
| ) | |
| if labels is not None: | |
| shift_logits = logits[..., :-1, :].reshape(-1, self.vocab_size) | |
| shift_labels = labels[..., 1:].reshape(-1) | |
| out["loss"] = F.cross_entropy(shift_logits, shift_labels, ignore_index=-100) | |
| return out | |
| # ---- v6: Mixture of Depths (dynamic layer skip) ---- | |
| if self.cfg.use_mod: | |
| mod_mask, mod_probs = self.mod(base_hidden) | |
| base_hidden = base_hidden * mod_mask | |
| # ---- coordination ---- | |
| gates, steps_dist, steps = self.coordinator(base_hidden) | |
| if not use_adaptive: | |
| gates = torch.ones_like(gates) * 0.9 | |
| steps = torch.full_like(steps, self.cfg.max_steps) | |
| # ---- System-2 reasoning loop ---- | |
| reasoned, scratch = self.reasoner(base_hidden, steps) | |
| # ---- multi-head CoT reasoning ---- | |
| reasoned = self.cot_reasoner(reasoned) | |
| # ---- subsystem gates ---- | |
| g_think, g_code, g_mem, g_scratch = gates[:, 0], gates[:, 1], gates[:, 2], gates[:, 3] | |
| thresh = self.cfg.adaptive_threshold | |
| # ---- code structure ---- | |
| if g_code.mean() >= thresh: | |
| reasoned = self.code_bias(reasoned, code_feat) | |
| # ---- depth conditioning ---- | |
| depth_emb, depth_dist = self.depth_module(reasoned.detach(), force_depth) | |
| refined = reasoned + depth_emb.unsqueeze(1) | |
| # ---- v6: multimodal fusion ---- | |
| refined = self.multimodal_fusion(refined, vision_embeds, audio_embeds, video_embeds) | |
| # ---- v6: document understanding ---- | |
| refined = self.document_module(refined, layout_info) | |
| # ---- v6: video understanding ---- | |
| refined = self.video_module(refined, frame_embeddings) | |
| # ---- sparse MoE (think gate) ---- | |
| use_moe = g_think.mean() >= thresh if use_adaptive else True | |
| if use_moe: | |
| moe_out, moe_aux = self.moe_block(refined) | |
| else: | |
| moe_out, moe_aux = torch.zeros_like(refined), torch.zeros((), device=refined.device) | |
| # ---- persistent memory read ---- | |
| use_mem = g_mem.mean() >= thresh if use_adaptive else True | |
| mem_out = self.memory.read(refined) if use_mem else torch.zeros_like(refined) | |
| # ---- scratchpad (cross-turn working memory) ---- | |
| use_scratch = g_scratch.mean() >= thresh if use_adaptive else True | |
| if use_scratch: | |
| refined = refined + self.scratchpad.read(refined) | |
| g = gates.mean(1) # (B,) mean gate, used to scale per batch | |
| refined = refined + g[:, None, None] * (moe_out + mem_out) | |
| # ---- lm head ---- | |
| lm_dtype = self.base_model_raw.lm_head.weight.dtype | |
| logits = self.base_model_raw.lm_head(refined.to(lm_dtype)) | |
| verifier_score = self.verifier(refined) | |
| tool_emb, tool_args = self.tool_use(refined) | |
| loss = None | |
| if labels is not None: | |
| shift_logits = logits[..., :-1, :].reshape(-1, self.vocab_size) | |
| shift_labels = labels[..., 1:].reshape(-1) | |
| ce = F.cross_entropy(shift_logits, shift_labels, ignore_index=-100) | |
| step_ent = -torch.sum(steps_dist * torch.log(steps_dist.clamp_min(1e-6)), dim=-1).mean() | |
| logp = -F.cross_entropy( | |
| shift_logits, shift_labels, reduction="none", ignore_index=-100 | |
| ).reshape(labels.shape[0], -1) | |
| mask = (labels[..., 1:] != -100).float() | |
| denom = mask.sum(1).clamp_min(1.0) | |
| seq_lik = (logp * mask).sum(1) / denom | |
| verifier_loss = F.mse_loss(verifier_score, seq_lik.detach()) | |
| # v6: MoD sparsity bonus (encourage more tokens to be skipped) | |
| mod_sparsity = mod_probs.mean() if self.cfg.use_mod else torch.tensor(0.0, device=refined.device) | |
| mod_loss = -torch.log(mod_sparsity.clamp_min(1e-6)).mean() * 0.01 | |
| loss = ( | |
| ce | |
| + self.cfg.moe_aux_weight * moe_aux | |
| + 0.01 * step_ent | |
| + self.cfg.verifier_weight * verifier_loss | |
| + mod_loss | |
| ) | |
| out = ModelOutput( | |
| logits=logits, | |
| hidden_states=base_hidden, | |
| refined=refined, | |
| gates=gates, | |
| steps_dist=steps_dist, | |
| steps=steps, | |
| depth_dist=depth_dist, | |
| verifier_score=verifier_score, | |
| tool_emb=tool_emb, | |
| tool_args=tool_args, | |
| loss=loss, | |
| ) | |
| self._last_refined = refined.detach() | |
| return out | |
| # ---- generation ---- | |
| def _greedy_step(self, inputs_embeds, attention_mask, skill_indices, code_feat, temperature, top_p): | |
| with torch.no_grad(): | |
| out = self.forward( | |
| inputs_embeds=inputs_embeds, | |
| attention_mask=attention_mask, | |
| skill_indices=skill_indices, | |
| code_feat=code_feat, | |
| use_adaptive=True, | |
| ) | |
| logits = out.logits[:, -1, :].float() / max(temperature, 1e-5) | |
| if top_p is not None and top_p < 1.0: | |
| sorted_logits, sorted_idx = torch.sort(logits, descending=True) | |
| cum = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1) | |
| mask = cum - F.softmax(sorted_logits, dim=-1) < top_p | |
| mask[:, 0] = True | |
| filtered = sorted_logits.clone() | |
| filtered[~mask] = float("-inf") | |
| logits = logits.scatter(-1, sorted_idx, filtered) | |
| probs = F.softmax(logits, dim=-1) | |
| return torch.multinomial(probs, num_samples=1) | |
| def generate( | |
| self, | |
| input_ids: Optional[torch.Tensor] = None, | |
| inputs_embeds: Optional[torch.Tensor] = None, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| skill_token_id: Optional[int] = None, | |
| code_feat: Optional[torch.Tensor] = None, | |
| max_new_tokens: int = 512, | |
| temperature: float = 0.7, | |
| top_p: float = 0.9, | |
| eos_token_id: Optional[int] = None, | |
| pad_token_id: Optional[int] = None, | |
| **kwargs, | |
| ): | |
| device = next(self.parameters()).device | |
| if inputs_embeds is None: | |
| inputs_embeds = self.base_model_raw.model.embed_tokens(input_ids.to(device)) | |
| skill_indices = None | |
| if skill_token_id is not None: | |
| skill_indices = torch.tensor([[skill_token_id]], dtype=torch.long, device=device) | |
| # skill injection is applied inside forward(), so we don't add it here | |
| if attention_mask is None: | |
| attention_mask = torch.ones(inputs_embeds.shape[:2], dtype=torch.long, device=device) | |
| gen = [] | |
| cur_emb = inputs_embeds | |
| attn = attention_mask | |
| for _ in range(max_new_tokens): | |
| nxt = self._greedy_step(cur_emb, attn, skill_indices, code_feat, temperature, top_p) | |
| gen.append(nxt) | |
| nxt_emb = self.base_model_raw.model.embed_tokens(nxt) | |
| cur_emb = torch.cat([cur_emb, nxt_emb], dim=1) | |
| attn = torch.cat([attn, torch.ones((attn.shape[0], 1), dtype=attn.dtype, device=device)], dim=1) | |
| if code_feat is not None: | |
| # extend code features with a neutral row to keep lengths aligned | |
| neutral = torch.zeros( | |
| (code_feat.shape[0], 1, code_feat.shape[-1]), | |
| dtype=code_feat.dtype, | |
| device=code_feat.device, | |
| ) | |
| neutral[..., 2] = 0.5 # neutral bracket balance | |
| code_feat = torch.cat([code_feat, neutral], dim=1) | |
| if eos_token_id is not None and (nxt == eos_token_id).all(): | |
| break | |
| gen_ids = torch.cat(gen, dim=1) | |
| if input_ids is not None: | |
| return torch.cat([input_ids.to(device), gen_ids], dim=1) | |
| return gen_ids | |
| def generate_best_of_n( | |
| self, | |
| input_ids: torch.Tensor, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| skill_token_id: Optional[int] = None, | |
| code_feat: Optional[torch.Tensor] = None, | |
| n: int = 4, | |
| max_new_tokens: int = 512, | |
| temperature: float = 0.9, | |
| top_p: float = 0.95, | |
| eos_token_id: Optional[int] = None, | |
| accept_threshold: Optional[float] = None, | |
| early_exit_margin: float = 0.01, | |
| **kwargs, | |
| ): | |
| """ | |
| Self-critique decoding. Generates up to n candidates and keeps the one | |
| the verifier scores highest. Early-exits (heuristic pruning) once a | |
| candidate clears `accept_threshold` and the marginal improvement over | |
| the previous best drops below `early_exit_margin`. Scores are | |
| normalized to [0,1] over the candidates seen so far so the threshold | |
| is stable across runs. | |
| """ | |
| eos_token_id = eos_token_id or self.tokenizer.eos_token_id | |
| best_ids, best_score = None, float("-inf") | |
| scores = [] | |
| for _ in range(n): | |
| cand = self.generate( | |
| input_ids=input_ids, | |
| attention_mask=attention_mask, | |
| skill_token_id=skill_token_id, | |
| code_feat=code_feat, | |
| max_new_tokens=max_new_tokens, | |
| temperature=temperature, | |
| top_p=top_p, | |
| eos_token_id=eos_token_id, | |
| ) | |
| with torch.no_grad(): | |
| out = self.forward( | |
| input_ids=cand, | |
| attention_mask=torch.ones_like(cand), | |
| skill_indices=( | |
| torch.tensor([[skill_token_id]], device=cand.device) | |
| if skill_token_id is not None | |
| else None | |
| ), | |
| code_feat=build_code_features(self.tokenizer, cand.cpu()).to(cand.device), | |
| use_adaptive=True, | |
| ) | |
| score = self.verifier(out.refined.detach()).item() | |
| scores.append(score) | |
| # min-max normalize against candidates generated so far | |
| lo, hi = min(scores), max(scores) | |
| norm = (score - lo) / (hi - lo) if hi > lo else 1.0 | |
| if score > best_score: | |
| best_score, best_ids = score, cand | |
| # heuristic search pruning: stop when good enough and no longer improving | |
| if ( | |
| accept_threshold is not None | |
| and norm >= accept_threshold | |
| and score <= best_score + early_exit_margin | |
| ): | |
| break | |
| return best_ids | |
| # ---- training helpers ---- | |
| def get_trainable_params(self) -> int: | |
| return sum(p.numel() for p in self.parameters() if p.requires_grad) | |
| def prune_experts(self): | |
| """Periodically prune underused MoE experts (call during training).""" | |
| self.moe_block.prune_experts() | |
| def state_dict(self, *args, **kwargs): | |
| sd = {} | |
| for name in ( | |
| "coordinator", "reasoner", "code_bias", "scratchpad", "verifier", | |
| "skill_module", "depth_module", "moe_block", "memory", | |
| "mod", "multimodal_fusion", "tool_use", "document_module", "video_module", | |
| "code_sandbox", "cot_reasoner", | |
| ): | |
| for k, v in getattr(self, name).state_dict().items(): | |
| sd[f"{name}.{k}"] = v | |
| for plugin_name, plugin in self._plugins.items(): | |
| for k, v in plugin.state_dict().items(): | |
| sd[f"plugin_{plugin_name}.{k}"] = v | |
| if self.base_model is not None: | |
| try: | |
| from peft import get_peft_model_state_dict | |
| sd.update(get_peft_model_state_dict(self.base_model)) | |
| except Exception as e: | |
| print(f"note: adapter state skipped ({e})") | |
| return sd | |
| def load_state_dict(self, sd, strict=True, assign=False): | |
| for name in ( | |
| "coordinator", "reasoner", "code_bias", "scratchpad", "verifier", | |
| "skill_module", "depth_module", "moe_block", "memory", | |
| "mod", "multimodal_fusion", "tool_use", "document_module", "video_module", | |
| "code_sandbox", "cot_reasoner", | |
| ): | |
| sub = {k[len(name) + 1:]: v for k, v in sd.items() if k.startswith(name + ".")} | |
| if sub: | |
| getattr(self, name).load_state_dict(sub) | |
| for plugin_name in self._plugins: | |
| prefix = f"plugin_{plugin_name}." | |
| sub = {k[len(prefix):]: v for k, v in sd.items() if k.startswith(prefix)} | |
| if sub: | |
| self._plugins[plugin_name].load_state_dict(sub) | |
| if self.base_model is not None: | |
| peft_sd = {k: v for k, v in sd.items() if k.startswith("base_model")} | |
| if peft_sd: | |
| from peft import set_peft_model_state_dict | |
| set_peft_model_state_dict(self.base_model, peft_sd) | |
| return {} | |
| def save_checkpoint(self, path: str): | |
| import os | |
| os.makedirs(path, exist_ok=True) | |
| torch.save( | |
| { | |
| "coordinator": self.coordinator.state_dict(), | |
| "reasoner": self.reasoner.state_dict(), | |
| "code_bias": self.code_bias.state_dict(), | |
| "scratchpad": self.scratchpad.state_dict(), | |
| "verifier": self.verifier.state_dict(), | |
| "skill_module": self.skill_module.state_dict(), | |
| "depth_module": self.depth_module.state_dict(), | |
| "moe_block": self.moe_block.state_dict(), | |
| "memory": self.memory.state_dict(), | |
| "mod": self.mod.state_dict(), | |
| "multimodal_fusion": self.multimodal_fusion.state_dict(), | |
| "tool_use": self.tool_use.state_dict(), | |
| "document_module": self.document_module.state_dict(), | |
| "video_module": self.video_module.state_dict(), | |
| "code_sandbox": self.code_sandbox.state_dict(), | |
| "cot_reasoner": self.cot_reasoner.state_dict(), | |
| **{f"plugin_{k}": v.state_dict() for k, v in self._plugins.items()}, | |
| "config": self.cfg, | |
| }, | |
| f"{path}/morph_components.pt", | |
| ) | |
| if self.base_model is not None: | |
| self.base_model.save_pretrained(f"{path}/base_lora") | |
| self.tokenizer.save_pretrained(path) | |
| print(f"Checkpoint saved to {path}") | |
| def load_checkpoint(self, path: str): | |
| import os | |
| from peft import PeftModel | |
| ckpt_path = f"{path}/morph_components.pt" | |
| if os.path.isfile(ckpt_path): | |
| ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False) | |
| for name in ( | |
| "coordinator", "reasoner", "code_bias", "scratchpad", "verifier", | |
| "skill_module", "depth_module", "moe_block", "memory", | |
| "mod", "multimodal_fusion", "tool_use", "document_module", "video_module", | |
| "code_sandbox", "cot_reasoner", | |
| ): | |
| if name in ckpt: | |
| getattr(self, name).load_state_dict(ckpt[name]) | |
| for plugin_name in self._plugins: | |
| key = f"plugin_{plugin_name}" | |
| if key in ckpt: | |
| self._plugins[plugin_name].load_state_dict(ckpt[key]) | |
| self.novel_trained = True | |
| else: | |
| trainer_ckpt = self._find_trainer_checkpoint(path) | |
| if trainer_ckpt: | |
| self._load_trainer_checkpoint(trainer_ckpt) | |
| self.novel_trained = True | |
| else: | |
| print(f"Note: no morph_components.pt at {path} - novel components use init weights") | |
| self.novel_trained = False | |
| lora_dir = f"{path}/base_lora" | |
| if os.path.isdir(lora_dir): | |
| self.base_model = PeftModel.from_pretrained(self.base_model_raw, lora_dir) | |
| print(f"LoRA adapter loaded from {lora_dir}") | |
| elif os.path.isfile(f"{path}/adapter_config.json"): | |
| self.base_model = PeftModel.from_pretrained(self.base_model_raw, path) | |
| print(f"LoRA adapter loaded from {path}") | |
| print(f"Checkpoint loaded from {path}") | |
| def _find_trainer_checkpoint(self, path: str): | |
| import glob | |
| candidates = sorted(glob.glob(f"{path}/checkpoint-*/model.safetensors")) | |
| return candidates[-1] if candidates else None | |
| def _load_trainer_checkpoint(self, ckpt_file: str): | |
| from safetensors import safe_open | |
| state_dict = {} | |
| with safe_open(ckpt_file, framework="pt") as f: | |
| for key in f.keys(): | |
| state_dict[key] = f.get_tensor(key) | |
| self.load_state_dict(state_dict) | |
| print(f"Loaded Trainer checkpoint from {ckpt_file}") |