Refactor JspaceAI codebase: remove main_multimodal.py and main_v2.py, add main_chat.py for interactive dialogue mode with online learning capabilities.
This commit is contained in:
@@ -2,7 +2,7 @@
|
||||
JspaceAI —— 全局工作空间 + J-space 广播的智慧系统
|
||||
|
||||
模块:
|
||||
1. core.py: 连续序列版(Expert + JSpaceWorkspace + JSpaceModel)
|
||||
1. core.py: 核心架构(Expert + JSpaceWorkspace + JSpaceModel,含 RK4/LayerNorm/异构专家)
|
||||
2. language_model.py: 语言版 + 自主进化
|
||||
3. jlens.py: J-lens 可解释性工具
|
||||
4. multimodal.py: 多模态(图像/音频/视频/文本)
|
||||
@@ -77,9 +77,6 @@ from .modules import (
|
||||
ExternalModule, SmallModelModule, KnowledgeBaseModule,
|
||||
ToolModule, ModuleDock,
|
||||
)
|
||||
from .core_v2 import JSpaceConfigV2, JSpaceModelV2
|
||||
from .distill_encoder import SmallModelEncoder
|
||||
from .distill_trainer import DistillationTrainer
|
||||
from .evolution import EvolutionTrainer
|
||||
|
||||
__all__ = [
|
||||
@@ -121,9 +118,6 @@ __all__ = [
|
||||
# 外挂模块系统(可热插拔)
|
||||
"ExternalModule", "SmallModelModule", "KnowledgeBaseModule",
|
||||
"ToolModule", "ModuleDock",
|
||||
# 改进版核心 v2
|
||||
"JSpaceConfigV2", "JSpaceModelV2",
|
||||
"SmallModelEncoder", "DistillationTrainer",
|
||||
# 自主进化
|
||||
"EvolutionTrainer",
|
||||
]
|
||||
|
||||
178
jspaceai/core.py
178
jspaceai/core.py
@@ -1,13 +1,13 @@
|
||||
"""
|
||||
核心架构:专家模块 + J-space 工作空间 + ODE 动力学
|
||||
|
||||
数学形式(每个 forward 时间步内做 Euler 积分若干子步):
|
||||
数学形式(每个 forward 时间步内做若干子步积分):
|
||||
|
||||
专家 i 的状态 m_i:
|
||||
dm_i/dt = -∇U_i(m_i) + J_i · w + P_i_in · x
|
||||
|
||||
U_i(m_i) = ½ ||m_i||² - ½ Σ_k softplus(a_ik · m_ik + b_ik)
|
||||
(多井势能:阻尼项 + softplus 形成的局部吸引子)
|
||||
U_i(m_i) = ½ ||m_i||² - ½ Σ_k softplus(a_ik · m_ik + b_ik) / sqrt(num_wells)
|
||||
(多井势能:阻尼项 + softplus 形成的局部吸引子,系数按井数归一化防梯度被 softplus 项主导)
|
||||
|
||||
工作空间 w:
|
||||
τ_w · dw/dt = -w + Σ_i α_i · P_i_out(m_i)
|
||||
@@ -16,6 +16,10 @@
|
||||
Jacobian 路由 J_i: 稀疏线性映射,每个专家只对 w 的少数维度敏感
|
||||
输出门控:当 ||w|| > θ 时触发输出 R(w)
|
||||
|
||||
支持两种 ODE 积分:Euler(简单可 backprop)和 RK4(更稳定,||w|| 量级显著提升)。
|
||||
可选 LayerNorm 防止 workspace 长期衰减。
|
||||
专家可携带模态特定编码器(异构专家),用于多模态场景。
|
||||
|
||||
所有参数都可 backprop 训练。学习目标是预测下一时刻的输入。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
@@ -23,7 +27,7 @@ from __future__ import annotations
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -40,6 +44,10 @@ class JSpaceConfig:
|
||||
output_threshold: float = 0.5 # 输出门控阈值(船舶涌出阈值)
|
||||
jacobian_sparsity: int = 8 # 每个 J_i 只保留前 k 大的连接
|
||||
noise_std: float = 0.01 # 内部噪声 ξ(t) 的标准差
|
||||
# —— v2 合并进来的改进字段(带默认值,向后兼容)——
|
||||
use_rk4: bool = True # 用 RK4 积分(比 Euler 稳定)
|
||||
use_layer_norm: bool = True # workspace 加 LayerNorm 防衰减
|
||||
well_coeff: float | None = None # softplus 项系数,None 则自动 1/sqrt(num_wells)
|
||||
|
||||
|
||||
class Expert(nn.Module):
|
||||
@@ -47,25 +55,36 @@ class Expert(nn.Module):
|
||||
单个专家模块。
|
||||
|
||||
状态:m_i ∈ R^{expert_dim}
|
||||
势能:U_i(m_i) = ½||m_i||² - ½ Σ_k softplus(a_k · m_i + b_k) · w_k
|
||||
势能:U_i(m_i) = ½||m_i||² - ½ · c · Σ_k softplus(a_k · m_i + b_k)
|
||||
- ½||m_i||² 是阻尼项(拉回原点)
|
||||
- softplus 项创造多个局部吸引子(多井势能 → 内部"思考")
|
||||
- c = well_coeff,默认 1/sqrt(num_wells),避免井数多时梯度被 softplus 主导
|
||||
动力学:dm_i/dt = -∇U_i(m_i) + J_i · w + P_in · x + ξ
|
||||
|
||||
J_i 是稀疏 Jacobian:从工作空间 w 路由信息进来。
|
||||
支持 RK4 积分(use_rk4=True)和 LayerNorm(use_layer_norm=True)。
|
||||
可选 encoder:把模态特定输入投影到 expert_dim(异构专家用)。
|
||||
"""
|
||||
|
||||
def __init__(self, expert_dim: int, workspace_dim: int, input_dim: int,
|
||||
num_wells: int, sparsity: int):
|
||||
num_wells: int, sparsity: int,
|
||||
use_rk4: bool = True, use_layer_norm: bool = True,
|
||||
well_coeff: float | None = None,
|
||||
encoder: nn.Module | None = None):
|
||||
super().__init__()
|
||||
self.expert_dim = expert_dim
|
||||
self.workspace_dim = workspace_dim
|
||||
self.num_wells = num_wells
|
||||
self.use_rk4 = use_rk4
|
||||
self.use_layer_norm = use_layer_norm
|
||||
|
||||
# softplus 项系数:默认按井数归一化,防止梯度被 softplus 项主导
|
||||
if well_coeff is None:
|
||||
well_coeff = 1.0 / (num_wells ** 0.5)
|
||||
self.well_coeff = well_coeff
|
||||
|
||||
# 势能景观参数:每个井是一个 softplus 形成的吸引子
|
||||
# U(m) = 0.5||m||^2 - 0.5 * sum_k softplus(a_k @ m + b_k)
|
||||
# ∇U(m) = m - 0.5 * sum_k sigmoid(a_k @ m + b_k) * a_k
|
||||
self.well_a = nn.Parameter(torch.randn(num_wells, expert_dim) * 0.3)
|
||||
self.well_a = nn.Parameter(torch.randn(num_wells, expert_dim) * 0.2)
|
||||
self.well_b = nn.Parameter(torch.zeros(num_wells))
|
||||
|
||||
# P_in: 输入投影 x -> m_i 的扰动
|
||||
@@ -75,42 +94,49 @@ class Expert(nn.Module):
|
||||
self.P_out = nn.Linear(expert_dim, workspace_dim, bias=False)
|
||||
|
||||
# J_i: 稀疏 Jacobian,从 w 路由信息到 m_i
|
||||
# 用 top-k 稀疏:训练时学习一个 full matrix,但只激活 top-k
|
||||
self.J_raw = nn.Parameter(torch.randn(expert_dim, workspace_dim) * 0.1)
|
||||
self.sparsity = sparsity
|
||||
|
||||
# 注意:sparsity 通过 forward 时 top-k 选择实现,可微性通过稀疏 mask 保留
|
||||
# 可选模态特定编码器(异构专家)
|
||||
self.encoder = encoder
|
||||
|
||||
# 可选 LayerNorm
|
||||
if use_layer_norm:
|
||||
self.m_norm = nn.LayerNorm(expert_dim)
|
||||
self.x_norm = nn.LayerNorm(expert_dim)
|
||||
|
||||
def get_sparse_J(self) -> torch.Tensor:
|
||||
"""获取稀疏化的 Jacobian:每行只保留 top-k 元素"""
|
||||
if self.sparsity >= self.workspace_dim:
|
||||
return self.J_raw
|
||||
# 对每行做 top-k(按绝对值)
|
||||
abs_J = self.J_raw.abs()
|
||||
topk_vals, topk_idx = abs_J.topk(self.sparsity, dim=-1)
|
||||
_, topk_idx = abs_J.topk(self.sparsity, dim=-1)
|
||||
mask = torch.zeros_like(self.J_raw)
|
||||
mask.scatter_(-1, topk_idx, 1.0)
|
||||
return self.J_raw * mask
|
||||
|
||||
def grad_potential(self, m: torch.Tensor) -> torch.Tensor:
|
||||
"""计算势能梯度 ∇U_i(m)
|
||||
U(m) = 0.5||m||^2 - 0.5 * sum_k softplus(a_k @ m + b_k)
|
||||
∇U(m) = m - 0.5 * sum_k sigmoid(a_k @ m + b_k) * a_k
|
||||
|
||||
U(m) = 0.5||m||^2 - 0.5 * c * sum_k softplus(a_k @ m + b_k)
|
||||
∇U(m) = m - 0.5 * c * sum_k sigmoid(a_k @ m + b_k) * a_k
|
||||
"""
|
||||
# m: (batch, expert_dim)
|
||||
# well_a: (num_wells, expert_dim)
|
||||
# a_k @ m: (batch, num_wells)
|
||||
am = F.linear(m, self.well_a, self.well_b) # (batch, num_wells)
|
||||
sig = torch.sigmoid(am) # (batch, num_wells)
|
||||
# sum_k sigmoid(...) * a_k: (batch, expert_dim)
|
||||
# well_a: (num_wells, expert_dim), sig: (batch, num_wells) -> (batch, 1, num_wells)
|
||||
# 用 matmul: sig @ well_a -> (batch, expert_dim)
|
||||
sig = torch.sigmoid(am) # (batch, num_wells)
|
||||
grad_wells = torch.matmul(sig, self.well_a) # (batch, expert_dim)
|
||||
return m - 0.5 * grad_wells
|
||||
return m - 0.5 * self.well_coeff * grad_wells
|
||||
|
||||
def deriv(self, m: torch.Tensor, w: torch.Tensor, x_feat: torch.Tensor) -> torch.Tensor:
|
||||
"""ODE 右端项 dm/dt = -∇U + J·w + P_in·x_feat"""
|
||||
J = self.get_sparse_J()
|
||||
w_proj = F.linear(w, J)
|
||||
x_proj = self.P_in(x_feat)
|
||||
grad_U = self.grad_potential(m)
|
||||
return -grad_U + w_proj + x_proj
|
||||
|
||||
def forward(self, m: torch.Tensor, w: torch.Tensor, x: torch.Tensor,
|
||||
dt: float, noise_std: float) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""一步 ODE 积分(Euler 法)
|
||||
"""一步 ODE 积分(Euler 或 RK4)
|
||||
|
||||
Args:
|
||||
m: (batch, expert_dim) 当前状态
|
||||
@@ -123,19 +149,30 @@ class Expert(nn.Module):
|
||||
m_next: (batch, expert_dim) 下一状态
|
||||
contribution: (batch, workspace_dim) 对工作空间的贡献(pre-attention)
|
||||
"""
|
||||
# 动力学: dm/dt = -∇U(m) + J·w + P_in·x + ξ
|
||||
J = self.get_sparse_J() # (expert_dim, workspace_dim)
|
||||
w_proj = F.linear(w, J) # (batch, expert_dim)
|
||||
x_proj = self.P_in(x) # (batch, expert_dim)
|
||||
grad_U = self.grad_potential(m) # (batch, expert_dim)
|
||||
# 可选编码器:把输入投影到专家内部空间
|
||||
if self.encoder is not None:
|
||||
x_feat = self.encoder(x)
|
||||
if self.use_layer_norm:
|
||||
x_feat = self.x_norm(x_feat)
|
||||
else:
|
||||
x_feat = x # P_in 会处理
|
||||
|
||||
noise = torch.randn_like(m) * noise_std if noise_std > 0 else 0.0
|
||||
if self.use_rk4:
|
||||
k1 = self.deriv(m, w, x_feat)
|
||||
k2 = self.deriv(m + 0.5 * dt * k1, w, x_feat)
|
||||
k3 = self.deriv(m + 0.5 * dt * k2, w, x_feat)
|
||||
k4 = self.deriv(m + dt * k3, w, x_feat)
|
||||
m_next = m + (dt / 6.0) * (k1 + 2 * k2 + 2 * k3 + k4)
|
||||
else:
|
||||
dm = self.deriv(m, w, x_feat)
|
||||
m_next = m + dt * dm
|
||||
|
||||
dm = -grad_U + w_proj + x_proj + noise
|
||||
m_next = m + dt * dm
|
||||
if noise_std > 0:
|
||||
m_next = m_next + torch.randn_like(m_next) * noise_std
|
||||
if self.use_layer_norm:
|
||||
m_next = self.m_norm(m_next)
|
||||
|
||||
# 对工作空间的贡献
|
||||
contribution = self.P_out(m_next) # (batch, workspace_dim)
|
||||
contribution = self.P_out(m_next)
|
||||
return m_next, contribution
|
||||
|
||||
|
||||
@@ -147,9 +184,11 @@ class JSpaceWorkspace(nn.Module):
|
||||
α_i = softmax(<q, P_i_out(m_i)>) q = MLP(x, w)
|
||||
|
||||
这个 α_i 是"注意力"——决定哪个专家的内容进入工作空间。
|
||||
可选 LayerNorm 防止 workspace 长期衰减。
|
||||
"""
|
||||
|
||||
def __init__(self, workspace_dim: int, input_dim: int, num_experts: int):
|
||||
def __init__(self, workspace_dim: int, input_dim: int, num_experts: int,
|
||||
use_layer_norm: bool = True):
|
||||
super().__init__()
|
||||
self.workspace_dim = workspace_dim
|
||||
|
||||
@@ -159,6 +198,9 @@ class JSpaceWorkspace(nn.Module):
|
||||
nn.Tanh(),
|
||||
nn.Linear(32, workspace_dim),
|
||||
)
|
||||
self.use_layer_norm = use_layer_norm
|
||||
if use_layer_norm:
|
||||
self.ln = nn.LayerNorm(workspace_dim)
|
||||
|
||||
def forward(self, w: torch.Tensor, x: torch.Tensor,
|
||||
contributions: torch.Tensor, dt: float,
|
||||
@@ -176,22 +218,15 @@ class JSpaceWorkspace(nn.Module):
|
||||
w_next: (batch, workspace_dim)
|
||||
alpha: (batch, num_experts) 注意力权重(可解释性用)
|
||||
"""
|
||||
# 生成 query
|
||||
q = self.query_gen(torch.cat([x, w], dim=-1)) # (batch, workspace_dim)
|
||||
|
||||
# 计算每个专家的注意力分数
|
||||
# contributions: (batch, num_experts, workspace_dim)
|
||||
# q: (batch, workspace_dim) -> (batch, 1, workspace_dim)
|
||||
scores = (contributions * q.unsqueeze(1)).sum(dim=-1) # (batch, num_experts)
|
||||
alpha = F.softmax(scores, dim=-1) # (batch, num_experts)
|
||||
alpha = F.softmax(scores, dim=-1)
|
||||
aggregated = (alpha.unsqueeze(-1) * contributions).sum(dim=1)
|
||||
|
||||
# 加权聚合
|
||||
# alpha: (batch, num_experts, 1) * contributions: (batch, num_experts, workspace_dim)
|
||||
aggregated = (alpha.unsqueeze(-1) * contributions).sum(dim=1) # (batch, workspace_dim)
|
||||
|
||||
# 动力学: τ_w · dw/dt = -w + aggregated
|
||||
dw = (-w + aggregated) / tau_w
|
||||
w_next = w + dt * dw
|
||||
if self.use_layer_norm:
|
||||
w_next = self.ln(w_next)
|
||||
return w_next, alpha
|
||||
|
||||
|
||||
@@ -219,6 +254,9 @@ class JSpaceModel(nn.Module):
|
||||
input_dim=config.input_dim,
|
||||
num_wells=config.num_wells,
|
||||
sparsity=config.jacobian_sparsity,
|
||||
use_rk4=config.use_rk4,
|
||||
use_layer_norm=config.use_layer_norm,
|
||||
well_coeff=config.well_coeff,
|
||||
)
|
||||
for _ in range(config.num_experts)
|
||||
])
|
||||
@@ -227,12 +265,13 @@ class JSpaceModel(nn.Module):
|
||||
workspace_dim=config.workspace_dim,
|
||||
input_dim=config.input_dim,
|
||||
num_experts=config.num_experts,
|
||||
use_layer_norm=config.use_layer_norm,
|
||||
)
|
||||
|
||||
# 输出门控:R(w) → action(这里 action = 预测的下一时刻输入)
|
||||
self.predictor = nn.Sequential(
|
||||
nn.Linear(config.workspace_dim, 32),
|
||||
nn.Tanh(),
|
||||
nn.GELU(),
|
||||
nn.Linear(32, config.input_dim),
|
||||
)
|
||||
|
||||
@@ -244,22 +283,25 @@ class JSpaceModel(nn.Module):
|
||||
for _ in range(self.config.num_experts)],
|
||||
}
|
||||
|
||||
def step(self, state: dict, x: torch.Tensor) -> tuple[dict, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
def step(self, state: dict, x: torch.Tensor,
|
||||
record_trajectory: bool = False) -> tuple[dict, torch.Tensor, torch.Tensor, torch.Tensor, list]:
|
||||
"""单时间步前向
|
||||
|
||||
Args:
|
||||
state: {'w': ..., 'm': [...]}
|
||||
x: (batch, input_dim) 输入
|
||||
record_trajectory: 是否记录每个子步的 w(J-lens 用)
|
||||
|
||||
Returns:
|
||||
new_state: 更新后的状态
|
||||
pred: (batch, input_dim) 预测的下一时刻输入
|
||||
alpha: (batch, num_experts) 注意力权重(可解释性)
|
||||
w_norm: (batch,) 工作空间范数(输出门控信号)
|
||||
new_state, pred, alpha, w_norm, w_trajectory
|
||||
"""
|
||||
w = state['w']
|
||||
ms = state['m']
|
||||
cfg = self.config
|
||||
w_trajectory = []
|
||||
|
||||
# ODE 子步积分
|
||||
for _ in range(cfg.ode_steps):
|
||||
# 1. 每个专家更新
|
||||
contributions = []
|
||||
new_ms = []
|
||||
for i, expert in enumerate(self.experts):
|
||||
@@ -269,31 +311,34 @@ class JSpaceModel(nn.Module):
|
||||
)
|
||||
new_ms.append(m_next)
|
||||
contributions.append(contrib)
|
||||
contributions = torch.stack(contributions, dim=1) # (batch, num_experts, workspace_dim)
|
||||
contributions = torch.stack(contributions, dim=1)
|
||||
|
||||
# 2. 工作空间更新
|
||||
w, alpha = self.workspace(
|
||||
w, x, contributions,
|
||||
dt=cfg.dt, tau_w=cfg.tau_w,
|
||||
)
|
||||
ms = new_ms
|
||||
|
||||
# 3. 输出:预测下一时刻输入(船舶涌出,但这里为了训练简化为每步都预测)
|
||||
if record_trajectory:
|
||||
w_trajectory.append(w.detach())
|
||||
|
||||
pred = self.predictor(w)
|
||||
w_norm = w.norm(dim=-1)
|
||||
|
||||
new_state = {'w': w, 'm': ms}
|
||||
return new_state, pred, alpha, w_norm
|
||||
return new_state, pred, alpha, w_norm, w_trajectory
|
||||
|
||||
def forward(self, xs: torch.Tensor, state: dict | None = None) -> tuple[torch.Tensor, dict]:
|
||||
def forward(self, xs: torch.Tensor, state: dict | None = None,
|
||||
record_trajectory: bool = False) -> tuple[torch.Tensor, dict]:
|
||||
"""
|
||||
Args:
|
||||
xs: (batch, T, input_dim) 输入序列
|
||||
state: 初始状态,None 则初始化
|
||||
record_trajectory: 是否记录每个时间步每个子步的 w
|
||||
|
||||
Returns:
|
||||
preds: (batch, T, input_dim) 每步对下一时刻的预测
|
||||
info: 包含注意力、w_norm 等可解释性信息
|
||||
info: 包含注意力、w_norm、可选 w_trajectory 等可解释性信息
|
||||
"""
|
||||
batch_size, T, _ = xs.shape
|
||||
device = xs.device
|
||||
@@ -304,17 +349,24 @@ class JSpaceModel(nn.Module):
|
||||
preds = []
|
||||
alphas = []
|
||||
w_norms = []
|
||||
w_traj_per_step = []
|
||||
for t in range(T):
|
||||
state, pred, alpha, w_norm = self.step(state, xs[:, t])
|
||||
state, pred, alpha, w_norm, w_traj = self.step(
|
||||
state, xs[:, t], record_trajectory=record_trajectory
|
||||
)
|
||||
preds.append(pred)
|
||||
alphas.append(alpha)
|
||||
w_norms.append(w_norm)
|
||||
if record_trajectory and w_traj:
|
||||
w_traj_per_step.append(torch.stack(w_traj, dim=1)) # (batch, n_substeps, workspace_dim)
|
||||
|
||||
preds = torch.stack(preds, dim=1) # (batch, T, input_dim)
|
||||
preds = torch.stack(preds, dim=1)
|
||||
info = {
|
||||
'alpha': torch.stack(alphas, dim=1), # (batch, T, num_experts)
|
||||
'w_norm': torch.stack(w_norms, dim=1), # (batch, T)
|
||||
'alpha': torch.stack(alphas, dim=1), # (batch, T, num_experts)
|
||||
'w_norm': torch.stack(w_norms, dim=1), # (batch, T)
|
||||
'final_w': state['w'],
|
||||
'final_m': state['m'],
|
||||
}
|
||||
if record_trajectory and w_traj_per_step:
|
||||
info['w_trajectory'] = torch.stack(w_traj_per_step, dim=1) # (batch, T, n_substeps, workspace_dim)
|
||||
return preds, info
|
||||
|
||||
@@ -1,208 +0,0 @@
|
||||
"""
|
||||
改进版核心架构 v2
|
||||
|
||||
四大改进:
|
||||
1. 异构专家:不同模态用不同架构
|
||||
2. workspace 扩容(256)+ LayerNorm 防衰减
|
||||
3. RK4 积分(比 Euler 稳定)
|
||||
4. 更大容量
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import torch, torch.nn as nn, torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass
|
||||
class JSpaceConfigV2:
|
||||
input_dim: int = 32
|
||||
workspace_dim: int = 256
|
||||
expert_dim: int = 64
|
||||
num_experts: int = 12
|
||||
num_wells: int = 8
|
||||
ode_steps: int = 4
|
||||
dt: float = 0.05
|
||||
tau_w: float = 1.0
|
||||
jacobian_sparsity: int = 32
|
||||
noise_std: float = 0.005
|
||||
use_rk4: bool = True
|
||||
use_layer_norm: bool = True
|
||||
|
||||
|
||||
class HeterogeneousExpert(nn.Module):
|
||||
"""异构专家基类——模态特定编码器 + ODE 动力学"""
|
||||
def __init__(self, expert_dim, workspace_dim, input_dim, num_wells, sparsity,
|
||||
use_rk4=True, use_layer_norm=True):
|
||||
super().__init__()
|
||||
self.expert_dim = expert_dim
|
||||
self.workspace_dim = workspace_dim
|
||||
self.use_rk4 = use_rk4
|
||||
|
||||
self.encoder = self._build_encoder(input_dim, expert_dim)
|
||||
self.well_a = nn.Parameter(torch.randn(num_wells, expert_dim) * 0.2)
|
||||
self.well_b = nn.Parameter(torch.zeros(num_wells))
|
||||
self.P_in = nn.Linear(expert_dim, expert_dim, bias=False)
|
||||
self.P_out = nn.Linear(expert_dim, workspace_dim, bias=False)
|
||||
self.J_raw = nn.Parameter(torch.randn(expert_dim, workspace_dim) * 0.05)
|
||||
self.sparsity = sparsity
|
||||
self.use_layer_norm = use_layer_norm
|
||||
if use_layer_norm:
|
||||
self.w_norm = nn.LayerNorm(expert_dim)
|
||||
self.m_norm = nn.LayerNorm(expert_dim)
|
||||
|
||||
def _build_encoder(self, input_dim, expert_dim):
|
||||
raise NotImplementedError
|
||||
|
||||
def get_sparse_J(self):
|
||||
if self.sparsity >= self.workspace_dim:
|
||||
return self.J_raw
|
||||
abs_J = self.J_raw.abs()
|
||||
_, topk_idx = abs_J.topk(self.sparsity, dim=-1)
|
||||
mask = torch.zeros_like(self.J_raw)
|
||||
mask.scatter_(-1, topk_idx, 1.0)
|
||||
return self.J_raw * mask
|
||||
|
||||
def grad_potential(self, m):
|
||||
am = F.linear(m, self.well_a, self.well_b)
|
||||
sig = torch.sigmoid(am)
|
||||
grad_wells = torch.matmul(sig, self.well_a)
|
||||
return m - 0.5 * grad_wells
|
||||
|
||||
def deriv(self, m, w, x_feat):
|
||||
J = self.get_sparse_J()
|
||||
w_proj = F.linear(w, J)
|
||||
x_proj = self.P_in(x_feat)
|
||||
grad_U = self.grad_potential(m)
|
||||
return -grad_U + w_proj + x_proj
|
||||
|
||||
def forward(self, m, w, x, dt, noise_std):
|
||||
x_feat = self.encoder(x)
|
||||
if self.use_layer_norm:
|
||||
x_feat = self.w_norm(x_feat)
|
||||
|
||||
if self.use_rk4:
|
||||
k1 = self.deriv(m, w, x_feat)
|
||||
k2 = self.deriv(m + 0.5*dt*k1, w, x_feat)
|
||||
k3 = self.deriv(m + 0.5*dt*k2, w, x_feat)
|
||||
k4 = self.deriv(m + dt*k3, w, x_feat)
|
||||
m_next = m + (dt/6.0)*(k1 + 2*k2 + 2*k3 + k4)
|
||||
else:
|
||||
dm = self.deriv(m, w, x_feat)
|
||||
m_next = m + dt * dm
|
||||
|
||||
if noise_std > 0:
|
||||
m_next = m_next + torch.randn_like(m_next) * noise_std
|
||||
if self.use_layer_norm:
|
||||
m_next = self.m_norm(m_next)
|
||||
|
||||
contribution = self.P_out(m_next)
|
||||
return m_next, contribution
|
||||
|
||||
|
||||
class VisualExpert(HeterogeneousExpert):
|
||||
def _build_encoder(self, input_dim, expert_dim):
|
||||
return nn.Sequential(
|
||||
nn.Linear(input_dim, 64), nn.ReLU(), nn.Linear(64, expert_dim))
|
||||
|
||||
class AudioExpert(HeterogeneousExpert):
|
||||
def _build_encoder(self, input_dim, expert_dim):
|
||||
return nn.Sequential(
|
||||
nn.Linear(input_dim, 64), nn.ReLU(), nn.Linear(64, expert_dim))
|
||||
|
||||
class LanguageExpert(HeterogeneousExpert):
|
||||
def _build_encoder(self, input_dim, expert_dim):
|
||||
return nn.Sequential(
|
||||
nn.Linear(input_dim, 64), nn.ReLU(), nn.Linear(64, expert_dim))
|
||||
|
||||
class CrossModalExpert(HeterogeneousExpert):
|
||||
def _build_encoder(self, input_dim, expert_dim):
|
||||
return nn.Sequential(
|
||||
nn.Linear(input_dim, 64), nn.ReLU(),
|
||||
nn.Linear(64, 64), nn.ReLU(), nn.Linear(64, expert_dim))
|
||||
|
||||
|
||||
class JSpaceWorkspaceV2(nn.Module):
|
||||
"""改进版 workspace——LayerNorm 防衰减"""
|
||||
def __init__(self, workspace_dim, input_dim, num_experts):
|
||||
super().__init__()
|
||||
self.workspace_dim = workspace_dim
|
||||
self.query_gen = nn.Sequential(
|
||||
nn.Linear(input_dim + workspace_dim, 128),
|
||||
nn.ReLU(), nn.Linear(128, workspace_dim))
|
||||
self.ln = nn.LayerNorm(workspace_dim)
|
||||
|
||||
def forward(self, w, x, contributions, dt, tau_w):
|
||||
q = self.query_gen(torch.cat([x, w], dim=-1))
|
||||
scores = (contributions * q.unsqueeze(1)).sum(dim=-1)
|
||||
alpha = F.softmax(scores, dim=-1)
|
||||
aggregated = (alpha.unsqueeze(-1) * contributions).sum(dim=1)
|
||||
dw = (-w + aggregated) / tau_w
|
||||
w_next = w + dt * dw
|
||||
w_next = self.ln(w_next)
|
||||
return w_next, alpha
|
||||
|
||||
|
||||
class JSpaceModelV2(nn.Module):
|
||||
"""改进版 JSpace 模型——异构专家 + 大 workspace + RK4"""
|
||||
def __init__(self, config: JSpaceConfigV2):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
expert_types = (
|
||||
[VisualExpert]*4 + [AudioExpert]*2 +
|
||||
[LanguageExpert]*2 + [CrossModalExpert]*4
|
||||
)[:config.num_experts]
|
||||
self.experts = nn.ModuleList([
|
||||
et(expert_dim=config.expert_dim,
|
||||
workspace_dim=config.workspace_dim,
|
||||
input_dim=config.input_dim,
|
||||
num_wells=config.num_wells,
|
||||
sparsity=config.jacobian_sparsity,
|
||||
use_rk4=config.use_rk4,
|
||||
use_layer_norm=config.use_layer_norm)
|
||||
for et in expert_types
|
||||
])
|
||||
self.expert_modality = (
|
||||
['visual']*2 + ['screen']*2 + ['audio']*2 +
|
||||
['text']*2 + ['mouse']*2 + ['cross']*2
|
||||
)[:config.num_experts]
|
||||
self.workspace = JSpaceWorkspaceV2(
|
||||
workspace_dim=config.workspace_dim,
|
||||
input_dim=config.input_dim,
|
||||
num_experts=config.num_experts,
|
||||
)
|
||||
|
||||
def init_state(self, batch_size, device):
|
||||
return {
|
||||
'w': torch.zeros(batch_size, self.config.workspace_dim, device=device),
|
||||
'm': [torch.zeros(batch_size, self.config.expert_dim, device=device)
|
||||
for _ in range(self.config.num_experts)],
|
||||
}
|
||||
|
||||
def step(self, state, x, record_trajectory=False):
|
||||
w, ms, cfg = state['w'], state['m'], self.config
|
||||
w_traj = []
|
||||
for _ in range(cfg.ode_steps):
|
||||
contributions, new_ms = [], []
|
||||
for i, expert in enumerate(self.experts):
|
||||
m_next, contrib = expert(ms[i], w, x, cfg.dt, cfg.noise_std)
|
||||
new_ms.append(m_next)
|
||||
contributions.append(contrib)
|
||||
contributions = torch.stack(contributions, dim=1)
|
||||
w, alpha = self.workspace(w, x, contributions, cfg.dt, cfg.tau_w)
|
||||
ms = new_ms
|
||||
if record_trajectory:
|
||||
w_traj.append(w.detach())
|
||||
return {'w': w, 'm': ms}, w_traj
|
||||
|
||||
def forward(self, xs, state=None, record_trajectory=False):
|
||||
B, T = xs.shape[0], xs.shape[1]
|
||||
device = xs.device
|
||||
if state is None:
|
||||
state = self.init_state(B, device)
|
||||
outputs, w_norms = [], []
|
||||
for t in range(T):
|
||||
state, _ = self.step(state, xs[:, t], record_trajectory)
|
||||
outputs.append(state['w'])
|
||||
w_norms.append(state['w'].norm(dim=-1))
|
||||
return torch.stack(outputs, dim=1), {
|
||||
'w_norm': torch.stack(w_norms, dim=1),
|
||||
}
|
||||
@@ -1,59 +0,0 @@
|
||||
"""小模型感知编码器——把小模型表征投影到 input_dim"""
|
||||
from __future__ import annotations
|
||||
import torch, torch.nn as nn, torch.nn.functional as F
|
||||
|
||||
|
||||
class SmallModelEncoder(nn.Module):
|
||||
def __init__(self, input_dim=32, model_name="gpt2", use_real_model=True):
|
||||
super().__init__()
|
||||
self.input_dim = input_dim
|
||||
self.model_name = model_name
|
||||
self._real_model = None
|
||||
self._tokenizer = None
|
||||
self._hidden_size = 128
|
||||
|
||||
if use_real_model:
|
||||
try:
|
||||
from transformers import AutoModel, AutoTokenizer
|
||||
self._tokenizer = AutoTokenizer.from_pretrained(model_name)
|
||||
self._model = AutoModel.from_pretrained(model_name)
|
||||
self._model.eval()
|
||||
self._hidden_size = self._model.config.hidden_size
|
||||
self._real_model = self._model
|
||||
if self._tokenizer.pad_token is None:
|
||||
self._tokenizer.pad_token = self._tokenizer.eos_token
|
||||
print(f" [编码器] 已加载 {model_name} (hidden={self._hidden_size})")
|
||||
except Exception as e:
|
||||
print(f" [编码器] 降级: {e}")
|
||||
|
||||
if self._real_model is not None:
|
||||
self.proj = nn.Linear(self._hidden_size, input_dim, bias=False)
|
||||
else:
|
||||
self.embed = nn.Embedding(1000, 64)
|
||||
self.proj = nn.Sequential(
|
||||
nn.Linear(64, 128), nn.ReLU(), nn.Linear(128, input_dim))
|
||||
|
||||
def forward(self, input_data):
|
||||
if self._real_model is not None:
|
||||
if isinstance(input_data, str):
|
||||
input_data = [input_data]
|
||||
if isinstance(input_data, list):
|
||||
with torch.no_grad():
|
||||
inputs = self._tokenizer(input_data, return_tensors="pt",
|
||||
truncation=True, max_length=64, padding=True)
|
||||
outputs = self._model(**inputs)
|
||||
hidden = outputs.last_hidden_state.mean(dim=1)
|
||||
return self.proj(hidden)
|
||||
if isinstance(input_data, torch.Tensor) and input_data.dtype == torch.long:
|
||||
with torch.no_grad():
|
||||
outputs = self._model(input_ids=input_data)
|
||||
hidden = outputs.last_hidden_state.mean(dim=1)
|
||||
return self.proj(hidden)
|
||||
if isinstance(input_data, torch.Tensor) and input_data.dtype == torch.long:
|
||||
emb = self.embed(input_data)
|
||||
if emb.dim() == 3: emb = emb.mean(dim=1)
|
||||
elif emb.dim() == 2: emb = emb.mean(dim=0, keepdim=True)
|
||||
return self.proj(emb)
|
||||
if isinstance(input_data, torch.Tensor):
|
||||
return self.proj(input_data)
|
||||
return torch.zeros(1, self.input_dim)
|
||||
@@ -1,82 +0,0 @@
|
||||
"""蒸馏训练器——把小模型理解能力迁移到 JspaceAI"""
|
||||
from __future__ import annotations
|
||||
import torch, torch.nn.functional as F, numpy as np
|
||||
|
||||
|
||||
class DistillationTrainer:
|
||||
"""蒸馏训练器
|
||||
|
||||
把小模型的理解能力迁移到 JspaceAI 编码器。
|
||||
训练后编码器离线也有理解能力,小模型可以断开。
|
||||
"""
|
||||
|
||||
def __init__(self, model_v2, encoder, texts, device='cpu'):
|
||||
self.model = model_v2.to(device)
|
||||
self.encoder = encoder.to(device)
|
||||
self.texts = texts
|
||||
self.device = device
|
||||
params = list(model_v2.parameters()) + list(encoder.proj.parameters())
|
||||
self.optimizer = torch.optim.Adam(params, lr=1e-3)
|
||||
self.history = []
|
||||
|
||||
def train_step(self, text_batch):
|
||||
# 1. 小模型表征(target)
|
||||
target = None
|
||||
if self.encoder._real_model is not None:
|
||||
with torch.no_grad():
|
||||
inputs = self.encoder._tokenizer(
|
||||
text_batch, return_tensors="pt",
|
||||
truncation=True, max_length=64, padding=True
|
||||
).to(self.device)
|
||||
outputs = self.encoder._model(**inputs)
|
||||
target = outputs.last_hidden_state.mean(dim=1)
|
||||
|
||||
# 2. 编码器输出
|
||||
encoded = self.encoder(text_batch)
|
||||
|
||||
# 3. workspace 处理
|
||||
xs = encoded.unsqueeze(1) # (batch, 1, input_dim)
|
||||
state = self.model.init_state(xs.shape[0], self.device)
|
||||
w_out, info = self.model(xs, state)
|
||||
|
||||
# 4. 蒸馏 loss
|
||||
if target is not None:
|
||||
target_proj = self.encoder.proj(target.to(self.device)).detach()
|
||||
distill_loss = F.mse_loss(encoded, target_proj)
|
||||
else:
|
||||
# 自监督
|
||||
if encoded.shape[0] > 1:
|
||||
distill_loss = F.mse_loss(encoded[:-1], encoded[1:].detach())
|
||||
else:
|
||||
distill_loss = torch.tensor(0.0, device=self.device)
|
||||
|
||||
# 5. workspace 稳定性(目标 ||w||≈1)
|
||||
w_norm = info['w_norm']
|
||||
stability_loss = F.mse_loss(w_norm, torch.ones_like(w_norm))
|
||||
|
||||
# 6. 总 loss
|
||||
total_loss = distill_loss + 0.1 * stability_loss
|
||||
|
||||
self.optimizer.zero_grad()
|
||||
total_loss.backward()
|
||||
torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
|
||||
self.optimizer.step()
|
||||
|
||||
self.history.append({
|
||||
'loss': total_loss.item(),
|
||||
'distill': distill_loss.item(),
|
||||
'stability': stability_loss.item(),
|
||||
'w_norm': w_norm.mean().item(),
|
||||
})
|
||||
return self.history[-1]
|
||||
|
||||
def train(self, n_steps=100, batch_size=4, verbose=True):
|
||||
for step in range(n_steps):
|
||||
batch = np.random.choice(self.texts, size=min(batch_size, len(self.texts)),
|
||||
replace=True).tolist()
|
||||
info = self.train_step(batch)
|
||||
if verbose and (step+1) % 20 == 0:
|
||||
print(f" step {step+1:4d} | loss {info['loss']:.4f} | "
|
||||
f"distill {info['distill']:.4f} | "
|
||||
f"||w|| {info['w_norm']:.3f}")
|
||||
return self.history
|
||||
Reference in New Issue
Block a user