Divergent Mechanisms of Antidepressant Efficacy: A Unified Computational Comparison of Synaptogenesis, Stabilization, and Tonic Inhibition in a Model of Depression.
The 10 matches
- [1] § Materials and methods ↔ 068A_antidep_modes_v2.ipynb, lines 47–131 · score 0.70 · rapid synaptogenesis, GABAergic, global, 50 %, receptor, 20 %
- [2] § Materials and methods ↔ 068A_antidep_modes_v2.ipynb, lines 47–131 · score 0.70 · bounded activation, GABAergic, enhanced tonic, accumulates, adaptation, baseline
- [3] § Materials and methods ↔ 068_antidep_modes.ipynb, lines 37–104 · score 0.69 · bounded activation, GABAergic, enhanced tonic, accumulates, adaptation, baseline
- [4] § Materials and methods ↔ 068A_antidep_modes_v2.ipynb, lines 1–37 · score 0.69 · iso dose comparison, cross mechanism, parameter sweeps, proxy, norm, L1
- [5] § Materials and methods ↔ 068A_antidep_modes_v2.ipynb, lines 174–251 · score 0.68 · feed forward network, ReLU, tonic inhibition, neuromodulatory, hidden, layers
- [6] § Materials and methods ↔ 068_antidep_modes.ipynb, lines 147–224 · score 0.68 · feed forward network, ReLU, tonic inhibition, neuromodulatory, hidden, layers
- [7] § Materials and methods ↔ 068_antidep_modes.ipynb, lines 37–104 · score 0.67 · rapid synaptogenesis, GABAergic, global, receptor, tanh, monoaminergic
- [8] § Materials and methods ↔ 068A_antidep_modes_v2.ipynb, lines 174–251 · score 0.55 · feed forward network, GABAergic, tonic inhibition, modules, hidden, layer
- [9] § Materials and methods ↔ 068_antidep_modes.ipynb, lines 147–224 · score 0.55 · feed forward network, GABAergic, tonic inhibition, modules, hidden, layer
- [10] § Materials and methods ↔ 068_antidep_modes.ipynb, lines 227–342 · score 0.51 · cross entropy loss, batch, threshold, magnitude, trained, weight
Paper
Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC
The paper is loaded when this pane is shown.
The authors' code
Jupyter notebook · 1,843 lines · 78 KB · no license · 5 matches
- # %% [markdown]
- # # Modes of Antidepressants
- # %%
- """
- ================================================================================
- MULTI-MECHANISM ANTIDEPRESSANT COMPARISON EXPERIMENT
- WITH ISO-DOSE FAIR COMPARISON PIPELINE
- ================================================================================
- This script runs only Experiment 4: comparing three antidepressant mechanisms:
- 1. KETAMINE-LIKE: Gradient-guided synaptogenesis
- 2. SSRI-LIKE: Gradual stabilization without structural changes
- 3. NEUROSTEROID-LIKE: Tonic inhibition enhancement
- All treatments start from identical pruned (depressed) network states.
- VERSION 2.0 ADDITIONS:
- - Iso-dose comparison pipeline for fair cross-mechanism comparison
- - L1/L2 weight change norms as mechanism-agnostic "dose" proxy
- - Synaptic turnover measurement
- - Parameter sweeps to match dose across treatments
- - Efficiency analysis: outcome per unit dose
- ================================================================================
- """
- import torch
- import torch.nn as nn
- import torch.optim as optim
- import numpy as np
- from torch.utils.data import DataLoader, TensorDataset
- from typing import Dict, Tuple, List, Optional, Any
- from dataclasses import dataclass
- import copy
- import warnings
- warnings.filterwarnings('ignore', category=UserWarning)
- # ============================================================================
- # REPRODUCIBILITY
- # ============================================================================
- SEED = 42
- torch.manual_seed(SEED)
- np.random.seed(SEED)
- DEVICE = torch.device('cpu')
- # ============================================================================
- # CONFIGURATION - Fixed version with proper syntax
- # ============================================================================
- CONFIG = {
- # Data generation
- 'n_train': 12000,
- 'n_test': 4000,
- 'n_clean_test': 2000,
- 'data_noise': 0.8,
- 'batch_size': 128,
- # Network architecture
- 'hidden_dims': [512, 512, 256],
- 'input_dim': 2,
- 'output_dim': 4,
- # Training hyperparameters
- 'baseline_epochs': 20,
- 'baseline_lr': 0.001,
- 'finetune_epochs': 15,
- 'finetune_lr': 0.0005,
- # Pruning parameters
- 'prune_sparsity': 0.95,
- # Regrowth parameters
- 'regrow_fraction': 0.5,
- 'regrow_init_scale': 0.03,
- 'gradient_accumulation_batches': 30,
- # Stress levels for evaluation
- 'extended_stress_levels': {
- 'none': 0.0,
- 'moderate': 0.5,
- 'high': 1.0,
- 'severe': 1.5,
- 'extreme': 2.5
- },
- # ========================================================================
- # MONOAMINERGIC (SSRI-LIKE) TREATMENT PARAMETERS
- # ========================================================================
- # Biological rationale: SSRIs increase synaptic serotonin, leading to
- # gradual receptor adaptations over weeks. No rapid synaptogenesis.
- # Network analog: Fixed sparsity, very low LR, gradual noise reduction.
- 'monoaminergic_epochs': 100,
- 'monoaminergic_lr': 1e-5,
- 'monoaminergic_initial_stress': 0.5,
- # ========================================================================
- # NEUROSTEROID (GABAergic) TREATMENT PARAMETERS
- # ========================================================================
- # Biological rationale: Neurosteroids enhance tonic GABA inhibition,
- # reducing network excitability rapidly (days, not weeks).
- # Network analog: Global activation damping, bounded activations.
- 'neurosteroid_inhibition_strength': 0.7,
- 'neurosteroid_use_tanh': True,
- 'neurosteroid_consolidation_epochs': 10,
- # ========================================================================
- # MULTI-MECHANISM COMPARISON PARAMETERS
- # ========================================================================
- 'comparison_ketamine_regrow': 0.5,
- 'comparison_ketamine_epochs': 15,
- 'comparison_ssri_epochs': 100,
- 'comparison_neurosteroid_strength': 0.7,
- 'comparison_neurosteroid_epochs': 10,
- # ========================================================================
- # ISO-DOSE COMPARISON PARAMETERS
- # ========================================================================
- 'iso_dose_norm_type': 'l1',
- 'iso_dose_target_doses': [0.005, 0.010, 0.020, 0.040],
- 'iso_dose_tolerance': 0.003,
- 'iso_dose_turnover_threshold': 0.10,
- # Parameter sweep ranges for iso-dose matching
- 'ketamine_regrow_sweep': [0.05, 0.10, 0.15, 0.20, 0.25, 0.30, 0.40, 0.50, 0.60, 0.70, 0.80],
- 'ssri_epochs_sweep': [10, 20, 30, 40, 50, 60, 80, 100, 120, 150, 200],
- 'ssri_lr_sweep': [1e-6, 5e-6, 1e-5, 2e-5, 5e-5],
- 'neurosteroid_strength_sweep': [0.50, 0.55, 0.60, 0.65, 0.70, 0.75, 0.80, 0.85, 0.90],
- # Relapse simulation parameters
- 'relapse_prune_fraction': 0.40,
- }
- # ============================================================================
- # DATA GENERATION
- # ============================================================================
- def generate_blobs(
- n_samples: int = 10000,
- noise: float = 0.8,
- seed: int = None
- ) -> Tuple[torch.Tensor, torch.Tensor]:
- """Generate 4-class Gaussian blob classification data."""
- if seed is not None:
- rng = np.random.RandomState(seed)
- else:
- rng = np.random.RandomState()
- centers = np.array([[-3, -3], [3, 3], [-3, 3], [3, -3]])
- labels = rng.randint(0, 4, n_samples)
- data = centers[labels] + rng.randn(n_samples, 2) * noise
- return (
- torch.tensor(data, dtype=torch.float32),
- torch.tensor(labels, dtype=torch.long)
- )
- def create_data_loaders() -> Tuple[DataLoader, DataLoader, DataLoader]:
- """Create train, test, and clean test data loaders."""
- train_data, train_labels = generate_blobs(CONFIG['n_train'], noise=CONFIG['data_noise'], seed=100)
- test_data, test_labels = generate_blobs(CONFIG['n_test'], noise=CONFIG['data_noise'], seed=200)
- clean_test_data, clean_test_labels = generate_blobs(CONFIG['n_clean_test'], noise=0.0, seed=300)
- train_loader = DataLoader(TensorDataset(train_data, train_labels), batch_size=CONFIG['batch_size'], shuffle=True)
- test_loader = DataLoader(TensorDataset(test_data, test_labels), batch_size=1000)
- clean_test_loader = DataLoader(TensorDataset(clean_test_data, clean_test_labels), batch_size=1000)
- return train_loader, test_loader, clean_test_loader
- train_loader, test_loader, clean_test_loader = create_data_loaders()
- # ============================================================================
- # NETWORK ARCHITECTURE
- # ============================================================================
- class StressAwareNetwork(nn.Module):
- """
- Feed-forward network with internal noise injection and GABAergic modulation.
- Supports three modulation mechanisms:
- - stress_level: Internal noise (neuromodulatory disruption)
- - inhibition_strength: Multiplicative damping (tonic GABA inhibition)
- - use_tanh: Bounded activation (shunting inhibition)
- """
- def __init__(self, hidden_dims: List[int] = None):
- super().__init__()
- if hidden_dims is None:
- hidden_dims = CONFIG['hidden_dims']
- self.fc1 = nn.Linear(CONFIG['input_dim'], hidden_dims[0])
- self.fc2 = nn.Linear(hidden_dims[0], hidden_dims[1])
- self.fc3 = nn.Linear(hidden_dims[1], hidden_dims[2])
- self.fc4 = nn.Linear(hidden_dims[2], CONFIG['output_dim'])
- self.relu = nn.ReLU()
- self.tanh = nn.Tanh()
- # Modulation parameters
- self.stress_level = 0.0 # Internal noise magnitude
- self.inhibition_strength = 1.0 # Multiplicative damping (1.0 = none)
- self.use_tanh = False # Use bounded activation
- self.weight_layers = ['fc1', 'fc2', 'fc3', 'fc4']
- def set_stress(self, level: float):
- """Set internal noise level for stress simulation."""
- self.stress_level = level
- def set_inhibition(self, strength: float, use_tanh: bool = False):
- """Set GABAergic tonic inhibition parameters."""
- self.inhibition_strength = strength
- self.use_tanh = use_tanh
- def reduce_stress_gradually(self, epoch: int, total_epochs: int,
- initial_stress: float = 0.5, final_stress: float = 0.0):
- """Linearly reduce internal stress over epochs (SSRI-like)."""
- progress = epoch / max(total_epochs - 1, 1)
- self.stress_level = initial_stress + progress * (final_stress - initial_stress)
- def forward(self, x: torch.Tensor) -> torch.Tensor:
- """Forward pass with noise injection and inhibitory modulation."""
- activation = self.tanh if self.use_tanh else self.relu
- # Layer 1
- h = activation(self.fc1(x))
- if self.stress_level > 0:
- h = h + torch.randn_like(h) * self.stress_level
- h = h * self.inhibition_strength
- # Layer 2
- h = activation(self.fc2(h))
- if self.stress_level > 0:
- h = h + torch.randn_like(h) * self.stress_level
- h = h * self.inhibition_strength
- # Layer 3
- h = activation(self.fc3(h))
- if self.stress_level > 0:
- h = h + torch.randn_like(h) * self.stress_level
- h = h * self.inhibition_strength
- # Output layer (no modulation)
- return self.fc4(h)
- def count_parameters(self) -> Tuple[int, int]:
- """Count total and non-zero parameters."""
- total = sum(p.numel() for p in self.parameters())
- nonzero = sum((p != 0).sum().item() for p in self.parameters())
- return total, nonzero
- # ============================================================================
- # PRUNING MANAGER
- # ============================================================================
- class PruningManager:
- """Manages structured pruning and gradient-guided regrowth."""
- def __init__(self, model: StressAwareNetwork):
- self.model = model
- self.masks = {}
- self.gradient_buffer = {}
- for name, param in model.named_parameters():
- if 'weight' in name and param.dim() >= 2:
- self.masks[name] = torch.ones_like(param, dtype=torch.float32)
- self.gradient_buffer[name] = torch.zeros_like(param)
- def prune_by_magnitude(self, sparsity: float, per_layer: bool = True) -> Dict[str, Dict]:
- """Prune weights by magnitude."""
- stats = {}
- for name, param in self.model.named_parameters():
- if name in self.masks:
- weights = param.data.abs()
- threshold = torch.quantile(weights.flatten(), sparsity)
- self.masks[name] = (weights >= threshold).float()
- param.data *= self.masks[name]
- kept = self.masks[name].sum().item()
- total = self.masks[name].numel()
- stats[name] = {'kept': int(kept), 'total': total, 'actual_sparsity': 1 - kept/total}
- return stats
- def _accumulate_gradients(self, num_batches: int = 30):
- """Accumulate gradient magnitudes at pruned positions."""
- model = self.model
- loss_fn = nn.CrossEntropyLoss()
- for name in self.gradient_buffer:
- self.gradient_buffer[name].zero_()
- model.train()
- original_stress = model.stress_level
- model.set_stress(0.0)
- batch_count = 0
- for x, y in train_loader:
- if batch_count >= num_batches:
- break
- x, y = x.to(DEVICE), y.to(DEVICE)
- loss = loss_fn(model(x), y)
- loss.backward()
- with torch.no_grad():
- for name, param in model.named_parameters():
- if name in self.masks:
- pruned_mask = (self.masks[name] == 0).float()
- self.gradient_buffer[name] += param.grad.abs() * pruned_mask
- model.zero_grad()
- batch_count += 1
- model.set_stress(original_stress)
- def gradient_guided_regrow(self, regrow_fraction: float,
- init_scale: float = None) -> Dict[str, Dict]:
- """Regrow pruned connections based on gradient importance."""
- if init_scale is None:
- init_scale = CONFIG['regrow_init_scale']
- self._accumulate_gradients(num_batches=CONFIG['gradient_accumulation_batches'])
- stats = {}
- for name, param in self.model.named_parameters():
- if name not in self.masks:
- continue
- mask = self.masks[name]
- pruned_positions = (mask == 0)
- num_pruned = pruned_positions.sum().item()
- if num_pruned == 0:
- stats[name] = {'regrown': 0, 'still_pruned': 0}
- continue
- gradient_scores = self.gradient_buffer[name][pruned_positions]
- num_regrow = max(1, int(regrow_fraction * num_pruned))
- num_regrow = min(num_regrow, gradient_scores.numel())
- _, top_indices = torch.topk(gradient_scores.flatten(), num_regrow)
- flat_pruned_indices = torch.where(pruned_positions.flatten())[0]
- regrow_flat_indices = flat_pruned_indices[top_indices]
- flat_mask = mask.flatten()
- flat_param = param.data.flatten()
- flat_mask[regrow_flat_indices] = 1.0
- flat_param[regrow_flat_indices] = torch.randn(num_regrow) * init_scale
- self.masks[name] = flat_mask.view_as(mask)
- param.data = flat_param.view_as(param)
- stats[name] = {'regrown': num_regrow, 'still_pruned': int(num_pruned - num_regrow)}
- return stats
- def apply_masks(self):
- """Re-apply masks to maintain sparsity."""
- with torch.no_grad():
- for name, param in self.model.named_parameters():
- if name in self.masks:
- param.data *= self.masks[name]
- def get_sparsity(self) -> float:
- """Calculate overall network sparsity."""
- total = sum(m.numel() for m in self.masks.values())
- zeros = sum((m == 0).sum().item() for m in self.masks.values())
- return zeros / total if total > 0 else 0.0
- def secondary_prune(self, fraction: float) -> Dict[str, Any]:
- """Simulate relapse by pruning a fraction of surviving weights."""
- stats = {}
- total_pruned = 0
- for name, param in self.model.named_parameters():
- if name not in self.masks:
- continue
- mask = self.masks[name]
- active_positions = (mask == 1)
- n_active = active_positions.sum().item()
- if n_active == 0:
- continue
- num_to_prune = int(fraction * n_active)
- if num_to_prune == 0:
- continue
- weights = param.data.abs()
- weights_active = weights.clone()
- weights_active[~active_positions] = float('inf')
- flat_weights = weights_active.flatten()
- threshold = torch.kthvalue(flat_weights, num_to_prune).values.item()
- prune_mask = (weights <= threshold) & active_positions
- mask[prune_mask] = 0
- param.data[prune_mask] = 0
- pruned_count = prune_mask.sum().item()
- total_pruned += pruned_count
- stats[name] = {
- 'pruned': pruned_count,
- 'remaining': n_active - pruned_count
- }
- new_sparsity = self.get_sparsity()
- return {
- 'total_pruned': total_pruned,
- 'new_sparsity': new_sparsity,
- 'layer_stats': stats
- }
- # ============================================================================
- # DOSING METRICS - MECHANISM-AGNOSTIC QUANTIFICATION
- # ============================================================================
- @dataclass
- class DoseMetrics:
- """Container for all dosing quantification metrics."""
- l1_norm: float = 0.0
- l2_norm: float = 0.0
- synaptic_turnover: float = 0.0
- sparsity_change: float = 0.0
- @property
- def primary_dose(self) -> float:
- """Primary dose metric (L1 norm by default)."""
- return self.l1_norm
- def compute_weight_change_norm(
- model_pre_state: Dict[str, torch.Tensor],
- model_post: nn.Module,
- norm_type: str = 'l1'
- ) -> float:
- """
- Compute total weight change magnitude as mechanism-agnostic dose proxy.
- Returns normalized dose (total change / total parameters).
- """
- delta = 0.0
- total_params = 0
- for name, param in model_post.named_parameters():
- if 'weight' in name and name in model_pre_state:
- diff = (param.data - model_pre_state[name]).abs()
- total_params += param.numel()
- if norm_type == 'l1':
- delta += diff.sum().item()
- elif norm_type == 'l2':
- delta += (diff ** 2).sum().item()
- if norm_type == 'l2':
- delta = delta ** 0.5
- return delta / total_params if total_params > 0 else 0.0
- def compute_synaptic_turnover(
- model_pre_state: Dict[str, torch.Tensor],
- model_post: nn.Module,
- threshold: float = 0.10
- ) -> float:
- """
- Compute fraction of synapses with significant weight changes.
- Captures "how many synapses were meaningfully modified" (> threshold relative change).
- """
- changed = 0
- total = 0
- for name, param in model_post.named_parameters():
- if 'weight' in name and name in model_pre_state:
- pre_weights = model_pre_state[name]
- relative_change = (param.data - pre_weights).abs() / (pre_weights.abs().clamp(min=1e-8))
- changed += (relative_change > threshold).sum().item()
- total += param.numel()
- return changed / total if total > 0 else 0.0
- def compute_sparsity_change(
- model_pre_state: Dict[str, torch.Tensor],
- model_post: nn.Module
- ) -> float:
- """Compute absolute change in network sparsity."""
- def get_sparsity(state_dict):
- total = 0
- zeros = 0
- for name, tensor in state_dict.items():
- if 'weight' in name:
- total += tensor.numel()
- zeros += (tensor.abs() < 1e-8).sum().item()
- return zeros / total if total > 0 else 0.0
- pre_sparsity = get_sparsity(model_pre_state)
- post_sparsity = get_sparsity({n: p.data for n, p in model_post.named_parameters()})
- return abs(post_sparsity - pre_sparsity)
- def compute_all_dose_metrics(
- model_pre_state: Dict[str, torch.Tensor],
- model_post: nn.Module,
- turnover_threshold: float = None
- ) -> DoseMetrics:
- """Compute all dose quantification metrics."""
- if turnover_threshold is None:
- turnover_threshold = CONFIG['iso_dose_turnover_threshold']
- return DoseMetrics(
- l1_norm=compute_weight_change_norm(model_pre_state, model_post, 'l1'),
- l2_norm=compute_weight_change_norm(model_pre_state, model_post, 'l2'),
- synaptic_turnover=compute_synaptic_turnover(model_pre_state, model_post, turnover_threshold),
- sparsity_change=compute_sparsity_change(model_pre_state, model_post)
- )
- # ============================================================================
- # TRAINING FUNCTIONS
- # ============================================================================
- def train(model: StressAwareNetwork, epochs: int = 15, lr: float = 0.001,
- pruning_manager: PruningManager = None, verbose: bool = False) -> List[float]:
- """Standard training loop."""
- optimizer = optim.Adam(model.parameters(), lr=lr)
- loss_fn = nn.CrossEntropyLoss()
- losses = []
- original_stress = model.stress_level
- model.set_stress(0.0)
- for epoch in range(epochs):
- model.train()
- epoch_loss = 0.0
- for x, y in train_loader:
- x, y = x.to(DEVICE), y.to(DEVICE)
- optimizer.zero_grad()
- loss = loss_fn(model(x), y)
- loss.backward()
- optimizer.step()
- if pruning_manager:
- pruning_manager.apply_masks()
- epoch_loss += loss.item()
- losses.append(epoch_loss / len(train_loader))
- if verbose:
- print(f" Epoch {epoch+1}/{epochs}, Loss: {losses[-1]:.4f}")
- model.set_stress(original_stress)
- return losses
- def train_with_stress_schedule(model: StressAwareNetwork, epochs: int, lr: float,
- initial_stress: float, final_stress: float = 0.0,
- pruning_manager: PruningManager = None,
- verbose: bool = False, print_interval: int = 20) -> List[float]:
- """Train with gradually reducing internal stress (SSRI-like)."""
- optimizer = optim.Adam(model.parameters(), lr=lr)
- loss_fn = nn.CrossEntropyLoss()
- losses = []
- for epoch in range(epochs):
- model.reduce_stress_gradually(epoch, epochs, initial_stress, final_stress)
- model.train()
- epoch_loss = 0.0
- for x, y in train_loader:
- x, y = x.to(DEVICE), y.to(DEVICE)
- optimizer.zero_grad()
- loss = loss_fn(model(x), y)
- loss.backward()
- optimizer.step()
- if pruning_manager:
- pruning_manager.apply_masks()
- epoch_loss += loss.item()
- losses.append(epoch_loss / len(train_loader))
- if verbose and (epoch + 1) % print_interval == 0:
- print(f" SSRI epoch {epoch+1}/{epochs}, stress: {model.stress_level:.3f}, loss: {losses[-1]:.4f}")
- model.set_stress(0.0)
- return losses
- # ============================================================================
- # EVALUATION FUNCTIONS
- # ============================================================================
- def evaluate(model: StressAwareNetwork, loader: DataLoader,
- input_noise: float = 0.0, internal_stress: float = 0.0) -> float:
- """Evaluate model accuracy."""
- model.eval()
- model.set_stress(internal_stress)
- correct, total = 0, 0
- with torch.no_grad():
- for x, y in loader:
- x, y = x.to(DEVICE), y.to(DEVICE)
- if input_noise > 0:
- x = x + torch.randn_like(x) * input_noise
- correct += (model(x).argmax(dim=1) == y).sum().item()
- total += y.size(0)
- model.set_stress(0.0)
- return 100.0 * correct / total
- def evaluate_with_neurosteroid(model: StressAwareNetwork, loader: DataLoader,
- input_noise: float = 0.0, internal_stress: float = 0.0) -> float:
- """Evaluate with neurosteroid modulation ACTIVE (inhibition settings preserved)."""
- model.eval()
- model.set_stress(internal_stress)
- correct, total = 0, 0
- with torch.no_grad():
- for x, y in loader:
- x, y = x.to(DEVICE), y.to(DEVICE)
- if input_noise > 0:
- x = x + torch.randn_like(x) * input_noise
- correct += (model(x).argmax(dim=1) == y).sum().item()
- total += y.size(0)
- model.set_stress(0.0)
- return 100.0 * correct / total
- # ============================================================================
- # TREATMENT PROTOCOLS
- # ============================================================================
- def ketamine_treatment(model: StressAwareNetwork, pruning_mgr: PruningManager,
- regrow_fraction: float = None, consolidation_epochs: int = None,
- verbose: bool = True) -> Dict:
- """
- KETAMINE-LIKE TREATMENT: Gradient-guided synaptogenesis.
- Biological model:
- - NMDA antagonism → BDNF release → mTOR activation → new spine formation
- - Activity-dependent targeting of new synapses
- - Brief consolidation to strengthen useful connections
- Key feature: ADDS NEW SYNAPSES (reduces sparsity)
- """
- if regrow_fraction is None:
- regrow_fraction = CONFIG['comparison_ketamine_regrow']
- if consolidation_epochs is None:
- consolidation_epochs = CONFIG['comparison_ketamine_epochs']
- if verbose:
- print(f"\n KETAMINE-LIKE TREATMENT:")
- print(f" Regrowth fraction: {regrow_fraction*100:.0f}%")
- print(f" Consolidation: {consolidation_epochs} epochs")
- print(f" Estimating gradient importance...")
- regrow_stats = pruning_mgr.gradient_guided_regrow(regrow_fraction=regrow_fraction)
- total_regrown = sum(s['regrown'] for s in regrow_stats.values())
- if verbose:
- print(f" Restored {total_regrown:,} synapses")
- print(f" Consolidating new synapses...")
- consolidation_losses = train(model, epochs=consolidation_epochs,
- lr=CONFIG['finetune_lr'], pruning_manager=pruning_mgr)
- final_sparsity = pruning_mgr.get_sparsity()
- if verbose:
- print(f" Final sparsity: {final_sparsity*100:.1f}%")
- return {'regrow_stats': regrow_stats, 'final_sparsity': final_sparsity}
- def ssri_treatment(model: StressAwareNetwork, pruning_mgr: PruningManager,
- epochs: int = None, learning_rate: float = None,
- initial_stress: float = None, verbose: bool = True,
- print_interval: int = 25) -> Dict:
- """
- SSRI-LIKE TREATMENT: Gradual stabilization without structural changes.
- Biological model:
- - Increased synaptic serotonin → gradual receptor adaptations
- - 5-HT1A autoreceptor desensitization over weeks
- - Improved signal-to-noise in existing circuits
- Key feature: NO NEW SYNAPSES (sparsity unchanged)
- """
- if epochs is None:
- epochs = CONFIG['monoaminergic_epochs']
- if learning_rate is None:
- learning_rate = CONFIG['monoaminergic_lr']
- if initial_stress is None:
- initial_stress = CONFIG['monoaminergic_initial_stress']
- if verbose:
- print(f"\n SSRI-LIKE TREATMENT:")
- print(f" Duration: {epochs} epochs (gradual)")
- print(f" Learning rate: {learning_rate} (very low)")
- print(f" Internal stress: {initial_stress} → 0.0")
- print(f" Note: NO structural changes (fixed sparsity)")
- initial_sparsity = pruning_mgr.get_sparsity()
- losses = train_with_stress_schedule(model, epochs=epochs, lr=learning_rate,
- initial_stress=initial_stress, final_stress=0.0,
- pruning_manager=pruning_mgr, verbose=verbose,
- print_interval=print_interval)
- final_sparsity = pruning_mgr.get_sparsity()
- if verbose:
- print(f" Final sparsity: {final_sparsity*100:.1f}% (unchanged)")
- return {'final_sparsity': final_sparsity, 'training_losses': losses}
- def neurosteroid_treatment(model: StressAwareNetwork, pruning_mgr: PruningManager,
- inhibition_strength: float = None, use_tanh: bool = None,
- consolidation_epochs: int = None, verbose: bool = True) -> Dict:
- """
- NEUROSTEROID-LIKE TREATMENT: Enhanced tonic inhibition.
- Biological model:
- - Enhanced extrasynaptic GABA-A receptor activation
- - Tonic (sustained) inhibition reduces network excitability
- - Rapid onset (days, not weeks)
- Key features:
- - NO NEW SYNAPSES (sparsity unchanged)
- - Works by DAMPING activity rather than building structure
- - Medication-dependent (effects reverse when stopped)
- """
- if inhibition_strength is None:
- inhibition_strength = CONFIG['neurosteroid_inhibition_strength']
- if use_tanh is None:
- use_tanh = CONFIG['neurosteroid_use_tanh']
- if consolidation_epochs is None:
- consolidation_epochs = CONFIG['neurosteroid_consolidation_epochs']
- if verbose:
- print(f"\n NEUROSTEROID-LIKE TREATMENT:")
- print(f" Inhibition strength: {inhibition_strength} ({(1-inhibition_strength)*100:.0f}% damping)")
- print(f" Bounded activation (tanh): {use_tanh}")
- print(f" Consolidation: {consolidation_epochs} epochs")
- print(f" Note: NO structural changes (fixed sparsity)")
- # Apply tonic inhibition modulation
- model.set_inhibition(inhibition_strength, use_tanh)
- if verbose:
- print(f" Applied tonic inhibition modulation...")
- print(f" Adapting to new activity dynamics...")
- consolidation_losses = train(model, epochs=consolidation_epochs,
- lr=CONFIG['finetune_lr'], pruning_manager=pruning_mgr)
- final_sparsity = pruning_mgr.get_sparsity()
- if verbose:
- print(f" Final sparsity: {final_sparsity*100:.1f}% (unchanged)")
- return {'final_sparsity': final_sparsity, 'inhibition_strength': inhibition_strength,
- 'use_tanh': use_tanh}
- # ============================================================================
- # ISO-DOSE PARAMETER SWEEP FUNCTIONS
- # ============================================================================
- def clone_model_and_manager(
- base_state_dict: Dict[str, torch.Tensor],
- base_masks: Dict[str, torch.Tensor]
- ) -> Tuple[StressAwareNetwork, PruningManager]:
- """Clone model and pruning manager from saved states."""
- model = StressAwareNetwork().to(DEVICE)
- model.load_state_dict({k: v.clone() for k, v in base_state_dict.items()})
- mgr = PruningManager(model)
- mgr.masks = {k: v.clone() for k, v in base_masks.items()}
- mgr.apply_masks()
- return model, mgr
- def run_ketamine_sweep(
- base_state_dict: Dict[str, torch.Tensor],
- base_masks: Dict[str, torch.Tensor],
- untreated_combined: float,
- regrow_fractions: List[float] = None,
- verbose: bool = False
- ) -> List[Dict[str, Any]]:
- """
- Sweep ketamine regrow_fraction and measure dose + outcomes.
- Returns list of results with dose metrics and performance outcomes.
- """
- if regrow_fractions is None:
- regrow_fractions = CONFIG['ketamine_regrow_sweep']
- results = []
- for regrow_frac in regrow_fractions:
- model, mgr = clone_model_and_manager(base_state_dict, base_masks)
- # Capture pre-treatment state
- pre_state = {n: p.data.clone() for n, p in model.named_parameters() if 'weight' in n}
- # Apply treatment
- ketamine_treatment(model, mgr, regrow_fraction=regrow_frac,
- consolidation_epochs=CONFIG['comparison_ketamine_epochs'],
- verbose=False)
- # Compute dose metrics
- dose_metrics = compute_all_dose_metrics(pre_state, model)
- # Evaluate acute performance
- acute_clean = evaluate(model, clean_test_loader, 0.0, 0.0)
- acute_standard = evaluate(model, test_loader, 0.0, 0.0)
- acute_combined = evaluate(model, test_loader, 1.0, 0.5)
- acute_extreme = evaluate(model, test_loader, 0.0, 2.5)
- # Relapse simulation
- pre_relapse_combined = acute_combined
- mgr.secondary_prune(fraction=CONFIG['relapse_prune_fraction'])
- post_relapse_combined = evaluate(model, test_loader, 1.0, 0.5)
- relapse_drop = pre_relapse_combined - post_relapse_combined
- improvement = acute_combined - untreated_combined
- results.append({
- 'treatment': 'ketamine',
- 'param_name': 'regrow_fraction',
- 'param_value': regrow_frac,
- 'dose': dose_metrics,
- 'sparsity': mgr.get_sparsity() * 100,
- 'acute_clean': acute_clean,
- 'acute_standard': acute_standard,
- 'acute_combined': acute_combined,
- 'acute_extreme': acute_extreme,
- 'improvement': improvement,
- 'relapse_drop': relapse_drop,
- 'post_relapse_combined': post_relapse_combined
- })
- if verbose:
- print(f" regrow_frac={regrow_frac:.2f}: dose={dose_metrics.l1_norm:.6f}, "
- f"combined={acute_combined:.1f}%, relapse_drop={relapse_drop:.1f}%")
- return results
- def run_ssri_sweep(
- base_state_dict: Dict[str, torch.Tensor],
- base_masks: Dict[str, torch.Tensor],
- untreated_combined: float,
- epochs_list: List[int] = None,
- lr_list: List[float] = None,
- verbose: bool = False
- ) -> List[Dict[str, Any]]:
- """
- Sweep SSRI epochs and learning rate, measure dose + outcomes.
- Returns list of results with dose metrics and performance outcomes.
- """
- if epochs_list is None:
- epochs_list = CONFIG['ssri_epochs_sweep']
- if lr_list is None:
- lr_list = [CONFIG['monoaminergic_lr']]
- results = []
- for epochs in epochs_list:
- for lr in lr_list:
- model, mgr = clone_model_and_manager(base_state_dict, base_masks)
- # Capture pre-treatment state
- pre_state = {n: p.data.clone() for n, p in model.named_parameters() if 'weight' in n}
- # Apply treatment
- ssri_treatment(model, mgr, epochs=epochs, learning_rate=lr,
- initial_stress=CONFIG['monoaminergic_initial_stress'],
- verbose=False)
- # Compute dose metrics
- dose_metrics = compute_all_dose_metrics(pre_state, model)
- # Evaluate acute performance
- acute_clean = evaluate(model, clean_test_loader, 0.0, 0.0)
- acute_standard = evaluate(model, test_loader, 0.0, 0.0)
- acute_combined = evaluate(model, test_loader, 1.0, 0.5)
- acute_extreme = evaluate(model, test_loader, 0.0, 2.5)
- # Relapse simulation
- pre_relapse_combined = acute_combined
- mgr.secondary_prune(fraction=CONFIG['relapse_prune_fraction'])
- post_relapse_combined = evaluate(model, test_loader, 1.0, 0.5)
- relapse_drop = pre_relapse_combined - post_relapse_combined
- improvement = acute_combined - untreated_combined
- results.append({
- 'treatment': 'ssri',
- 'param_name': 'epochs',
- 'param_value': epochs,
- 'lr': lr,
- 'dose': dose_metrics,
- 'sparsity': mgr.get_sparsity() * 100,
- 'acute_clean': acute_clean,
- 'acute_standard': acute_standard,
- 'acute_combined': acute_combined,
- 'acute_extreme': acute_extreme,
- 'improvement': improvement,
- 'relapse_drop': relapse_drop,
- 'post_relapse_combined': post_relapse_combined
- })
- if verbose:
- print(f" epochs={epochs}, lr={lr:.0e}: dose={dose_metrics.l1_norm:.6f}, "
- f"combined={acute_combined:.1f}%, relapse_drop={relapse_drop:.1f}%")
- return results
- def run_neurosteroid_sweep(
- base_state_dict: Dict[str, torch.Tensor],
- base_masks: Dict[str, torch.Tensor],
- untreated_combined: float,
- strength_list: List[float] = None,
- verbose: bool = False
- ) -> List[Dict[str, Any]]:
- """
- Sweep neurosteroid inhibition strength, measure dose + outcomes.
- Returns list of results with dose metrics and performance outcomes.
- """
- if strength_list is None:
- strength_list = CONFIG['neurosteroid_strength_sweep']
- results = []
- for strength in strength_list:
- model, mgr = clone_model_and_manager(base_state_dict, base_masks)
- # Capture pre-treatment state
- pre_state = {n: p.data.clone() for n, p in model.named_parameters() if 'weight' in n}
- # Apply treatment
- neurosteroid_treatment(model, mgr, inhibition_strength=strength,
- use_tanh=CONFIG['neurosteroid_use_tanh'],
- consolidation_epochs=CONFIG['neurosteroid_consolidation_epochs'],
- verbose=False)
- # Compute dose metrics
- dose_metrics = compute_all_dose_metrics(pre_state, model)
- # Evaluate acute performance WITH modulation active
- acute_clean = evaluate_with_neurosteroid(model, clean_test_loader, 0.0, 0.0)
- acute_standard = evaluate_with_neurosteroid(model, test_loader, 0.0, 0.0)
- acute_combined = evaluate_with_neurosteroid(model, test_loader, 1.0, 0.5)
- acute_extreme = evaluate_with_neurosteroid(model, test_loader, 0.0, 2.5)
- # Evaluate OFF medication
- model.set_inhibition(1.0, False)
- off_med_combined = evaluate(model, test_loader, 1.0, 0.5)
- off_med_extreme = evaluate(model, test_loader, 0.0, 2.5)
- off_med_reversal = off_med_combined - acute_combined
- # Restore modulation for relapse test
- model.set_inhibition(strength, CONFIG['neurosteroid_use_tanh'])
- # Relapse simulation (with modulation active)
- pre_relapse_combined = acute_combined
- mgr.secondary_prune(fraction=CONFIG['relapse_prune_fraction'])
- post_relapse_combined = evaluate_with_neurosteroid(model, test_loader, 1.0, 0.5)
- relapse_drop = pre_relapse_combined - post_relapse_combined
- improvement = acute_combined - untreated_combined
- results.append({
- 'treatment': 'neurosteroid',
- 'param_name': 'strength',
- 'param_value': strength,
- 'dose': dose_metrics,
- 'sparsity': mgr.get_sparsity() * 100,
- 'acute_clean': acute_clean,
- 'acute_standard': acute_standard,
- 'acute_combined': acute_combined,
- 'acute_extreme': acute_extreme,
- 'improvement': improvement,
- 'off_med_combined': off_med_combined,
- 'off_med_extreme': off_med_extreme,
- 'off_med_reversal': off_med_reversal,
- 'relapse_drop': relapse_drop,
- 'post_relapse_combined': post_relapse_combined
- })
- if verbose:
- print(f" strength={strength:.2f}: dose={dose_metrics.l1_norm:.6f}, "
- f"combined={acute_combined:.1f}%, off_med={off_med_combined:.1f}%, "
- f"relapse_drop={relapse_drop:.1f}%")
- return results
- def find_iso_dose_match(
- sweep_results: List[Dict[str, Any]],
- target_dose: float,
- tolerance: float = None
- ) -> Optional[Dict[str, Any]]:
- """
- Find parameter configuration closest to target dose.
- Returns the result with dose closest to target, or None if no results.
- """
- if tolerance is None:
- tolerance = CONFIG['iso_dose_tolerance']
- if not sweep_results:
- return None
- best_match = None
- best_diff = float('inf')
- for result in sweep_results:
- dose = result['dose'].l1_norm
- diff = abs(dose - target_dose)
- if diff < best_diff:
- best_diff = diff
- best_match = result
- return best_match
- def compute_efficiency(improvement: float, dose: float) -> float:
- """Compute treatment efficiency: improvement per unit dose."""
- return improvement / (dose + 1e-8)
- # ============================================================================
- # ISO-DOSE COMPARISON EXPERIMENT
- # ============================================================================
- def run_iso_dose_comparison_experiment(
- base_state_dict: Dict[str, torch.Tensor],
- base_masks: Dict[str, torch.Tensor],
- untreated_results: Dict[str, Any]
- ) -> Dict[str, Any]:
- """
- Run iso-dose comparison across all three treatment mechanisms.
- Sweeps parameters for each treatment, measures dose (L1 weight change norm),
- and compares outcomes at matched dose levels.
- """
- print("\n" + "=" * 80)
- print(" ISO-DOSE FAIR COMPARISON EXPERIMENT")
- print("=" * 80)
- print("""
- ISO-DOSE COMPARISON METHODOLOGY:
- ┌────────────────────────────────────────────────────────────────────────────┐
- │ DOSE METRIC: L1 weight change norm (normalized per parameter) │
- │ │
- │ This provides a mechanism-agnostic measure of "how much the network │
- │ changed" regardless of whether changes came from: │
- │ - Structural regrowth (ketamine) │
- │ - Gradual weight refinement (SSRI) │
- │ - Consolidation under modulation (neurosteroid) │
- │ │
- │ EFFICIENCY = (Performance Improvement) / (Dose) │
- │ Higher efficiency = better outcome per unit of network change │
- └────────────────────────────────────────────────────────────────────────────┘
- """)
- untreated_combined = untreated_results['combined']
- results = {
- 'sweeps': {},
- 'iso_dose_comparisons': {},
- 'efficiency_analysis': {}
- }
- # ========================================================================
- # PHASE 1: Parameter Sweeps
- # ========================================================================
- print("-" * 70)
- print(" PHASE 1: Parameter Sweeps (Measuring Dose-Response)")
- print("-" * 70)
- print("\n [KETAMINE] Sweeping regrow_fraction...")
- print(f" Testing {len(CONFIG['ketamine_regrow_sweep'])} configurations...")
- ketamine_results = run_ketamine_sweep(
- base_state_dict, base_masks, untreated_combined, verbose=True
- )
- results['sweeps']['ketamine'] = ketamine_results
- print(f" Completed {len(ketamine_results)} configurations")
- print("\n [SSRI] Sweeping epochs...")
- print(f" Testing {len(CONFIG['ssri_epochs_sweep'])} configurations...")
- ssri_results = run_ssri_sweep(
- base_state_dict, base_masks, untreated_combined, verbose=True
- )
- results['sweeps']['ssri'] = ssri_results
- print(f" Completed {len(ssri_results)} configurations")
- print("\n [NEUROSTEROID] Sweeping inhibition strength...")
- print(f" Testing {len(CONFIG['neurosteroid_strength_sweep'])} configurations...")
- neurosteroid_results = run_neurosteroid_sweep(
- base_state_dict, base_masks, untreated_combined, verbose=True
- )
- results['sweeps']['neurosteroid'] = neurosteroid_results
- print(f" Completed {len(neurosteroid_results)} configurations")
- # ========================================================================
- # PHASE 2: Dose-Response Analysis
- # ========================================================================
- print("\n" + "-" * 70)
- print(" PHASE 2: Dose-Response Curves")
- print("-" * 70)
- print("\n KETAMINE DOSE-RESPONSE:")
- print(f" {'regrow_frac':>12} {'L1 Dose':>12} {'L2 Dose':>12} {'Turnover':>10} "
- f"{'Combined':>10} {'Improve':>10} {'Relapse':>10}")
- print(" " + "-" * 88)
- for r in ketamine_results:
- print(f" {r['param_value']:>12.2f} {r['dose'].l1_norm:>12.6f} {r['dose'].l2_norm:>12.6f} "
- f"{r['dose'].synaptic_turnover:>10.4f} {r['acute_combined']:>9.1f}% "
- f"{r['improvement']:>+9.1f}% {r['relapse_drop']:>9.1f}%")
- print("\n SSRI DOSE-RESPONSE:")
- print(f" {'epochs':>12} {'L1 Dose':>12} {'L2 Dose':>12} {'Turnover':>10} "
- f"{'Combined':>10} {'Improve':>10} {'Relapse':>10}")
- print(" " + "-" * 88)
- for r in ssri_results:
- print(f" {r['param_value']:>12} {r['dose'].l1_norm:>12.6f} {r['dose'].l2_norm:>12.6f} "
- f"{r['dose'].synaptic_turnover:>10.4f} {r['acute_combined']:>9.1f}% "
- f"{r['improvement']:>+9.1f}% {r['relapse_drop']:>9.1f}%")
- print("\n NEUROSTEROID DOSE-RESPONSE:")
- print(f" {'strength':>12} {'L1 Dose':>12} {'L2 Dose':>12} {'Turnover':>10} "
- f"{'Combined':>10} {'Off-Med':>10} {'Relapse':>10}")
- print(" " + "-" * 88)
- for r in neurosteroid_results:
- print(f" {r['param_value']:>12.2f} {r['dose'].l1_norm:>12.6f} {r['dose'].l2_norm:>12.6f} "
- f"{r['dose'].synaptic_turnover:>10.4f} {r['acute_combined']:>9.1f}% "
- f"{r['off_med_combined']:>9.1f}% {r['relapse_drop']:>9.1f}%")
- # ========================================================================
- # PHASE 3: Determine Dose Range for Iso-Dose Matching
- # ========================================================================
- print("\n" + "-" * 70)
- print(" PHASE 3: Iso-Dose Matching")
- print("-" * 70)
- all_doses = []
- for r in ketamine_results:
- all_doses.append(('ketamine', r['param_value'], r['dose'].l1_norm))
- for r in ssri_results:
- all_doses.append(('ssri', r['param_value'], r['dose'].l1_norm))
- for r in neurosteroid_results:
- all_doses.append(('neurosteroid', r['param_value'], r['dose'].l1_norm))
- dose_values = [d[2] for d in all_doses]
- dose_min, dose_max = min(dose_values), max(dose_values)
- print(f"\n Observed dose range across all treatments:")
- print(f" Minimum dose: {dose_min:.6f}")
- print(f" Maximum dose: {dose_max:.6f}")
- # Find overlapping dose range
- ket_doses = [r['dose'].l1_norm for r in ketamine_results]
- ssri_doses = [r['dose'].l1_norm for r in ssri_results]
- neuro_doses = [r['dose'].l1_norm for r in neurosteroid_results]
- overlap_min = max(min(ket_doses), min(ssri_doses), min(neuro_doses))
- overlap_max = min(max(ket_doses), max(ssri_doses), max(neuro_doses))
- print(f"\n Overlapping dose range (for fair comparison):")
- print(f" Overlap minimum: {overlap_min:.6f}")
- print(f" Overlap maximum: {overlap_max:.6f}")
- # Generate target doses within overlap
- if overlap_max > overlap_min:
- target_doses = np.linspace(overlap_min, overlap_max, 5).tolist()
- else:
- # If no overlap, use percentiles of all doses
- target_doses = [np.percentile(dose_values, p) for p in [20, 40, 60, 80]]
- print(f"\n Target doses for iso-dose comparison:")
- for i, td in enumerate(target_doses):
- print(f" Level {i+1}: {td:.6f}")
- # ========================================================================
- # PHASE 4: Iso-Dose Matched Comparisons
- # ========================================================================
- print("\n" + "-" * 70)
- print(" PHASE 4: Iso-Dose Matched Comparisons")
- print("-" * 70)
- for i, target_dose in enumerate(target_doses):
- print(f"\n ISO-DOSE LEVEL {i+1}: Target Dose = {target_dose:.6f}")
- print(" " + "=" * 75)
- ket_match = find_iso_dose_match(ketamine_results, target_dose)
- ssri_match = find_iso_dose_match(ssri_results, target_dose)
- neuro_match = find_iso_dose_match(neurosteroid_results, target_dose)
- iso_comparison = {
- 'target_dose': target_dose,
- 'ketamine': ket_match,
- 'ssri': ssri_match,
- 'neurosteroid': neuro_match
- }
- results['iso_dose_comparisons'][f'level_{i+1}'] = iso_comparison
- print(f"\n {'Treatment':<15} {'Parameter':<18} {'Actual Dose':>12} {'Dose Diff':>12} "
- f"{'Combined':>10} {'Improve':>10} {'Relapse':>10}")
- print(" " + "-" * 95)
- for name, match in [('Ketamine', ket_match), ('SSRI', ssri_match), ('Neurosteroid', neuro_match)]:
- if match:
- param_str = f"{match['param_name']}={match['param_value']}"
- actual_dose = match['dose'].l1_norm
- dose_diff = actual_dose - target_dose
- print(f" {name:<15} {param_str:<18} {actual_dose:>12.6f} {dose_diff:>+12.6f} "
- f"{match['acute_combined']:>9.1f}% {match['improvement']:>+9.1f}% "
- f"{match['relapse_drop']:>9.1f}%")
- else:
- print(f" {name:<15} {'N/A':<18} {'N/A':>12} {'N/A':>12} {'N/A':>10} {'N/A':>10} {'N/A':>10}")
- # Determine best treatment at this dose level
- valid_matches = [(n, m) for n, m in [('Ketamine', ket_match), ('SSRI', ssri_match),
- ('Neurosteroid', neuro_match)] if m]
- if valid_matches:
- best_combined = max(valid_matches, key=lambda x: x[1]['acute_combined'])
- best_relapse = min(valid_matches, key=lambda x: x[1]['relapse_drop'])
- print(f"\n At this dose level:")
- print(f" Best acute performance: {best_combined[0]} ({best_combined[1]['acute_combined']:.1f}%)")
- print(f" Best relapse resistance: {best_relapse[0]} (drop: {best_relapse[1]['relapse_drop']:.1f}%)")
- # ========================================================================
- # PHASE 5: Efficiency Analysis
- # ========================================================================
- print("\n" + "-" * 70)
- print(" PHASE 5: Treatment Efficiency Analysis")
- print("-" * 70)
- print("\n EFFICIENCY = (Performance Improvement) / (L1 Dose)")
- print(" Higher efficiency = better outcome per unit of network change")
- print("\n KETAMINE EFFICIENCY:")
- print(f" {'regrow_frac':>12} {'Dose':>12} {'Improvement':>12} {'Efficiency':>12}")
- print(" " + "-" * 52)
- ket_efficiencies = []
- for r in ketamine_results:
- eff = compute_efficiency(r['improvement'], r['dose'].l1_norm)
- ket_efficiencies.append(eff)
- print(f" {r['param_value']:>12.2f} {r['dose'].l1_norm:>12.6f} {r['improvement']:>+11.1f}% {eff:>12.2f}")
- print("\n SSRI EFFICIENCY:")
- print(f" {'epochs':>12} {'Dose':>12} {'Improvement':>12} {'Efficiency':>12}")
- print(" " + "-" * 52)
- ssri_efficiencies = []
- for r in ssri_results:
- eff = compute_efficiency(r['improvement'], r['dose'].l1_norm)
- ssri_efficiencies.append(eff)
- print(f" {r['param_value']:>12} {r['dose'].l1_norm:>12.6f} {r['improvement']:>+11.1f}% {eff:>12.2f}")
- print("\n NEUROSTEROID EFFICIENCY:")
- print(f" {'strength':>12} {'Dose':>12} {'Improvement':>12} {'Efficiency':>12}")
- print(" " + "-" * 52)
- neuro_efficiencies = []
- for r in neurosteroid_results:
- eff = compute_efficiency(r['improvement'], r['dose'].l1_norm)
- neuro_efficiencies.append(eff)
- print(f" {r['param_value']:>12.2f} {r['dose'].l1_norm:>12.6f} {r['improvement']:>+11.1f}% {eff:>12.2f}")
- # Store efficiency results
- results['efficiency_analysis'] = {
- 'ketamine': {
- 'max_efficiency': max(ket_efficiencies) if ket_efficiencies else 0,
- 'mean_efficiency': np.mean(ket_efficiencies) if ket_efficiencies else 0,
- 'best_config': ketamine_results[np.argmax(ket_efficiencies)]['param_value'] if ket_efficiencies else None
- },
- 'ssri': {
- 'max_efficiency': max(ssri_efficiencies) if ssri_efficiencies else 0,
- 'mean_efficiency': np.mean(ssri_efficiencies) if ssri_efficiencies else 0,
- 'best_config': ssri_results[np.argmax(ssri_efficiencies)]['param_value'] if ssri_efficiencies else None
- },
- 'neurosteroid': {
- 'max_efficiency': max(neuro_efficiencies) if neuro_efficiencies else 0,
- 'mean_efficiency': np.mean(neuro_efficiencies) if neuro_efficiencies else 0,
- 'best_config': neurosteroid_results[np.argmax(neuro_efficiencies)]['param_value'] if neuro_efficiencies else None
- }
- }
- # ========================================================================
- # PHASE 6: Summary Statistics
- # ========================================================================
- print("\n" + "-" * 70)
- print(" PHASE 6: Summary Statistics")
- print("-" * 70)
- def compute_sweep_stats(sweep_results, treatment_name):
- doses = [r['dose'].l1_norm for r in sweep_results]
- turnovers = [r['dose'].synaptic_turnover for r in sweep_results]
- improvements = [r['improvement'] for r in sweep_results]
- relapse_drops = [r['relapse_drop'] for r in sweep_results]
- efficiencies = [compute_efficiency(r['improvement'], r['dose'].l1_norm) for r in sweep_results]
- return {
- 'treatment': treatment_name,
- 'n_configs': len(sweep_results),
- 'dose_range': (min(doses), max(doses)),
- 'dose_mean': np.mean(doses),
- 'turnover_range': (min(turnovers), max(turnovers)),
- 'best_improvement': max(improvements),
- 'worst_improvement': min(improvements),
- 'best_relapse': min(relapse_drops),
- 'worst_relapse': max(relapse_drops),
- 'max_efficiency': max(efficiencies),
- 'mean_efficiency': np.mean(efficiencies)
- }
- ket_stats = compute_sweep_stats(ketamine_results, 'Ketamine')
- ssri_stats = compute_sweep_stats(ssri_results, 'SSRI')
- neuro_stats = compute_sweep_stats(neurosteroid_results, 'Neurosteroid')
- results['summary_stats'] = {
- 'ketamine': ket_stats,
- 'ssri': ssri_stats,
- 'neurosteroid': neuro_stats
- }
- print("\n TREATMENT SUMMARY ACROSS ALL CONFIGURATIONS:")
- print(f"\n {'Metric':<25} {'Ketamine':>18} {'SSRI':>18} {'Neurosteroid':>18}")
- print(" " + "-" * 80)
- print(f" {'Configurations tested':<25} {ket_stats['n_configs']:>18} "
- f"{ssri_stats['n_configs']:>18} {neuro_stats['n_configs']:>18}")
- print(f" {'Dose range (L1)':<25} {ket_stats['dose_range'][0]:.4f}-{ket_stats['dose_range'][1]:.4f}"
- f" {ssri_stats['dose_range'][0]:.4f}-{ssri_stats['dose_range'][1]:.4f}"
- f" {neuro_stats['dose_range'][0]:.4f}-{neuro_stats['dose_range'][1]:.4f}")
- print(f" {'Mean dose':<25} {ket_stats['dose_mean']:>18.6f} "
- f"{ssri_stats['dose_mean']:>18.6f} {neuro_stats['dose_mean']:>18.6f}")
- print(f" {'Best improvement':<25} {ket_stats['best_improvement']:>+17.1f}% "
- f"{ssri_stats['best_improvement']:>+17.1f}% {neuro_stats['best_improvement']:>+17.1f}%")
- print(f" {'Best relapse resistance':<25} {ket_stats['best_relapse']:>17.1f}% "
- f"{ssri_stats['best_relapse']:>17.1f}% {neuro_stats['best_relapse']:>17.1f}%")
- print(f" {'Max efficiency':<25} {ket_stats['max_efficiency']:>18.2f} "
- f"{ssri_stats['max_efficiency']:>18.2f} {neuro_stats['max_efficiency']:>18.2f}")
- print(f" {'Mean efficiency':<25} {ket_stats['mean_efficiency']:>18.2f} "
- f"{ssri_stats['mean_efficiency']:>18.2f} {neuro_stats['mean_efficiency']:>18.2f}")
- # ========================================================================
- # PHASE 7: Detailed Iso-Dose Comparison Tables
- # ========================================================================
- print("\n" + "=" * 80)
- print(" DETAILED ISO-DOSE COMPARISON RESULTS")
- print("=" * 80)
- for level_key, comparison in results['iso_dose_comparisons'].items():
- target_dose = comparison['target_dose']
- print(f"\n {level_key.upper().replace('_', ' ')}: TARGET DOSE = {target_dose:.6f}")
- print(" " + "=" * 75)
- print(f"\n {'Treatment':<15} {'Parameter':<20} {'Actual Dose':>12} {'Turnover':>10} "
- f"{'ΔSparsity':>10}")
- print(" " + "-" * 75)
- for treatment_name in ['ketamine', 'ssri', 'neurosteroid']:
- match = comparison.get(treatment_name)
- if match:
- param_str = f"{match['param_name']}={match['param_value']}"
- d = match['dose']
- print(f" {treatment_name.capitalize():<15} {param_str:<20} {d.l1_norm:>12.6f} "
- f"{d.synaptic_turnover:>10.4f} {d.sparsity_change:>10.4f}")
- else:
- print(f" {treatment_name.capitalize():<15} {'N/A':<20} {'N/A':>12} {'N/A':>10} {'N/A':>10}")
- print(f"\n {'Treatment':<15} {'Clean':>10} {'Standard':>10} {'Combined':>10} "
- f"{'Extreme':>10} {'Improve':>10} {'Relapse':>10}")
- print(" " + "-" * 85)
- for treatment_name in ['ketamine', 'ssri', 'neurosteroid']:
- match = comparison.get(treatment_name)
- if match:
- print(f" {treatment_name.capitalize():<15} {match['acute_clean']:>9.1f}% "
- f"{match['acute_standard']:>9.1f}% {match['acute_combined']:>9.1f}% "
- f"{match['acute_extreme']:>9.1f}% {match['improvement']:>+9.1f}% "
- f"{match['relapse_drop']:>9.1f}%")
- else:
- print(f" {treatment_name.capitalize():<15} {'N/A':>10} {'N/A':>10} {'N/A':>10} "
- f"{'N/A':>10} {'N/A':>10} {'N/A':>10}")
- # Neurosteroid-specific: off-medication performance
- neuro_match = comparison.get('neurosteroid')
- if neuro_match and 'off_med_combined' in neuro_match:
- print(f"\n NEUROSTEROID OFF-MEDICATION:")
- print(f" Combined stress (on medication): {neuro_match['acute_combined']:.1f}%")
- print(f" Combined stress (off medication): {neuro_match['off_med_combined']:.1f}%")
- print(f" Performance reversal: {neuro_match['off_med_reversal']:+.1f}%")
- # ========================================================================
- # PHASE 8: Cross-Treatment Comparison at Each Dose Level
- # ========================================================================
- print("\n" + "=" * 80)
- print(" CROSS-TREATMENT RANKINGS AT EACH DOSE LEVEL")
- print("=" * 80)
- for level_key, comparison in results['iso_dose_comparisons'].items():
- target_dose = comparison['target_dose']
- print(f"\n {level_key.upper().replace('_', ' ')}: Dose ≈ {target_dose:.6f}")
- print(" " + "-" * 60)
- valid_treatments = []
- for treatment_name in ['ketamine', 'ssri', 'neurosteroid']:
- match = comparison.get(treatment_name)
- if match:
- valid_treatments.append({
- 'name': treatment_name.capitalize(),
- 'combined': match['acute_combined'],
- 'improvement': match['improvement'],
- 'relapse_drop': match['relapse_drop'],
- 'efficiency': compute_efficiency(match['improvement'], match['dose'].l1_norm)
- })
- if valid_treatments:
- # Rank by combined performance
- by_combined = sorted(valid_treatments, key=lambda x: x['combined'], reverse=True)
- print(f"\n Ranking by Combined Stress Performance:")
- for rank, t in enumerate(by_combined, 1):
- print(f" {rank}. {t['name']:<15} {t['combined']:.1f}%")
- # Rank by relapse resistance (lower is better)
- by_relapse = sorted(valid_treatments, key=lambda x: x['relapse_drop'])
- print(f"\n Ranking by Relapse Resistance (lower drop = better):")
- for rank, t in enumerate(by_relapse, 1):
- print(f" {rank}. {t['name']:<15} {t['relapse_drop']:.1f}% drop")
- # Rank by efficiency
- by_efficiency = sorted(valid_treatments, key=lambda x: x['efficiency'], reverse=True)
- print(f"\n Ranking by Efficiency (improvement per dose):")
- for rank, t in enumerate(by_efficiency, 1):
- print(f" {rank}. {t['name']:<15} {t['efficiency']:.2f}")
- return results
- # ============================================================================
- # MAIN EXPERIMENT
- # ============================================================================
- def run_multi_mechanism_experiment() -> Dict[str, Dict]:
- """
- Compare ketamine, SSRI, and neurosteroid treatment mechanisms.
- All treatments start from identical 95% sparse (depressed) networks.
- """
- print("\n" + "="*80)
- print(" MULTI-MECHANISM ANTIDEPRESSANT COMPARISON EXPERIMENT")
- print("="*80)
- print("""
- COMPARING THREE ANTIDEPRESSANT MECHANISMS:
- ┌─────────────────┬─────────────────────────────────────────────────────────┐
- │ Mechanism │ Key Feature │
- ├─────────────────┼─────────────────────────────────────────────────────────┤
- │ Ketamine │ Gradient-guided synaptogenesis (↑ density) │
- │ SSRI │ Gradual noise reduction (stabilizes existing weights) │
- │ Neurosteroid │ Tonic inhibition (damps activity, bounds firing) │
- └─────────────────┴─────────────────────────────────────────────────────────┘
- All treatments start from identical 95% sparse (depressed) networks.
- """)
- # ========================================================================
- # PREPARE BASE PRUNED MODEL
- # ========================================================================
- print("-"*70)
- print(" Preparing shared pruned baseline...")
- print("-"*70)
- base_model = StressAwareNetwork().to(DEVICE)
- print(f" Training full network ({CONFIG['baseline_epochs']} epochs)...")
- train(base_model, epochs=CONFIG['baseline_epochs'], lr=CONFIG['baseline_lr'])
- base_pruning_mgr = PruningManager(base_model)
- base_pruning_mgr.prune_by_magnitude(sparsity=CONFIG['prune_sparsity'], per_layer=True)
- initial_sparsity = base_pruning_mgr.get_sparsity()
- print(f" Pruned to {initial_sparsity*100:.1f}% sparse")
- # Evaluate untreated state
- print("\n UNTREATED PRUNED STATE:")
- untreated_results = {'sparsity': initial_sparsity * 100}
- untreated_results['clean'] = evaluate(base_model, clean_test_loader, 0.0, 0.0)
- untreated_results['standard'] = evaluate(base_model, test_loader, 0.0, 0.0)
- for stress_name, stress_level in CONFIG['extended_stress_levels'].items():
- untreated_results[f'stress_{stress_name}'] = evaluate(base_model, test_loader, 0.0, stress_level)
- untreated_results['combined'] = evaluate(base_model, test_loader, 1.0, 0.5)
- print(f" Clean: {untreated_results['clean']:.1f}%")
- print(f" Standard: {untreated_results['standard']:.1f}%")
- print(f" Combined stress: {untreated_results['combined']:.1f}%")
- print(f" Extreme stress: {untreated_results['stress_extreme']:.1f}%")
- # Save state for cloning
- base_state_dict = {k: v.clone() for k, v in base_model.state_dict().items()}
- base_masks = {k: v.clone() for k, v in base_pruning_mgr.masks.items()}
- results = {'untreated': untreated_results}
- # ========================================================================
- # TREATMENT 1: KETAMINE-LIKE
- # ========================================================================
- print("\n" + "="*70)
- print(" TREATMENT 1: KETAMINE-LIKE (Synaptogenesis)")
- print("="*70)
- ketamine_model = StressAwareNetwork().to(DEVICE)
- ketamine_model.load_state_dict(base_state_dict)
- ketamine_mgr = PruningManager(ketamine_model)
- ketamine_mgr.masks = {k: v.clone() for k, v in base_masks.items()}
- ketamine_mgr.apply_masks()
- ketamine_stats = ketamine_treatment(ketamine_model, ketamine_mgr,
- regrow_fraction=CONFIG['comparison_ketamine_regrow'],
- consolidation_epochs=CONFIG['comparison_ketamine_epochs'])
- print("\n POST-TREATMENT EVALUATION:")
- ketamine_results = {'sparsity': ketamine_mgr.get_sparsity() * 100}
- ketamine_results['clean'] = evaluate(ketamine_model, clean_test_loader, 0.0, 0.0)
- ketamine_results['standard'] = evaluate(ketamine_model, test_loader, 0.0, 0.0)
- for stress_name, stress_level in CONFIG['extended_stress_levels'].items():
- ketamine_results[f'stress_{stress_name}'] = evaluate(ketamine_model, test_loader, 0.0, stress_level)
- ketamine_results['combined'] = evaluate(ketamine_model, test_loader, 1.0, 0.5)
- print(f" Clean: {ketamine_results['clean']:.1f}%")
- print(f" Standard: {ketamine_results['standard']:.1f}%")
- print(f" Combined stress: {ketamine_results['combined']:.1f}%")
- print(f" Extreme stress: {ketamine_results['stress_extreme']:.1f}%")
- # Relapse simulation
- print("\n RELAPSE SIMULATION:")
- pre_relapse = ketamine_results['combined']
- pre_sparsity = ketamine_mgr.get_sparsity()
- target_sparsity = min(pre_sparsity + (1 - pre_sparsity) * 0.40, 0.99)
- ketamine_mgr.prune_by_magnitude(sparsity=target_sparsity, per_layer=True)
- ketamine_mgr.apply_masks()
- post_relapse = evaluate(ketamine_model, test_loader, 1.0, 0.5)
- ketamine_results['relapse_drop'] = pre_relapse - post_relapse
- print(f" Combined: {pre_relapse:.1f}% → {post_relapse:.1f}% (drop: {ketamine_results['relapse_drop']:.1f}%)")
- results['ketamine'] = ketamine_results
- # ========================================================================
- # TREATMENT 2: SSRI-LIKE
- # ========================================================================
- print("\n" + "="*70)
- print(" TREATMENT 2: SSRI-LIKE (Gradual Stabilization)")
- print("="*70)
- ssri_model = StressAwareNetwork().to(DEVICE)
- ssri_model.load_state_dict(base_state_dict)
- ssri_mgr = PruningManager(ssri_model)
- ssri_mgr.masks = {k: v.clone() for k, v in base_masks.items()}
- ssri_mgr.apply_masks()
- ssri_stats = ssri_treatment(ssri_model, ssri_mgr,
- epochs=CONFIG['comparison_ssri_epochs'],
- learning_rate=CONFIG['monoaminergic_lr'],
- initial_stress=CONFIG['monoaminergic_initial_stress'],
- print_interval=25)
- print("\n POST-TREATMENT EVALUATION:")
- ssri_results = {'sparsity': ssri_mgr.get_sparsity() * 100}
- ssri_results['clean'] = evaluate(ssri_model, clean_test_loader, 0.0, 0.0)
- ssri_results['standard'] = evaluate(ssri_model, test_loader, 0.0, 0.0)
- for stress_name, stress_level in CONFIG['extended_stress_levels'].items():
- ssri_results[f'stress_{stress_name}'] = evaluate(ssri_model, test_loader, 0.0, stress_level)
- ssri_results['combined'] = evaluate(ssri_model, test_loader, 1.0, 0.5)
- print(f" Clean: {ssri_results['clean']:.1f}%")
- print(f" Standard: {ssri_results['standard']:.1f}%")
- print(f" Combined stress: {ssri_results['combined']:.1f}%")
- print(f" Extreme stress: {ssri_results['stress_extreme']:.1f}%")
- # Relapse simulation
- print("\n RELAPSE SIMULATION:")
- pre_relapse = ssri_results['combined']
- pre_sparsity = ssri_mgr.get_sparsity()
- target_sparsity = min(pre_sparsity + (1 - pre_sparsity) * 0.40, 0.99)
- ssri_mgr.prune_by_magnitude(sparsity=target_sparsity, per_layer=True)
- ssri_mgr.apply_masks()
- post_relapse = evaluate(ssri_model, test_loader, 1.0, 0.5)
- ssri_results['relapse_drop'] = pre_relapse - post_relapse
- print(f" Combined: {pre_relapse:.1f}% → {post_relapse:.1f}% (drop: {ssri_results['relapse_drop']:.1f}%)")
- results['ssri'] = ssri_results
- # ========================================================================
- # TREATMENT 3: NEUROSTEROID-LIKE
- # ========================================================================
- print("\n" + "="*70)
- print(" TREATMENT 3: NEUROSTEROID-LIKE (Tonic Inhibition)")
- print("="*70)
- neuro_model = StressAwareNetwork().to(DEVICE)
- neuro_model.load_state_dict(base_state_dict)
- neuro_mgr = PruningManager(neuro_model)
- neuro_mgr.masks = {k: v.clone() for k, v in base_masks.items()}
- neuro_mgr.apply_masks()
- neuro_stats = neurosteroid_treatment(neuro_model, neuro_mgr,
- inhibition_strength=CONFIG['neurosteroid_inhibition_strength'],
- use_tanh=CONFIG['neurosteroid_use_tanh'],
- consolidation_epochs=CONFIG['neurosteroid_consolidation_epochs'])
- # Evaluate WITH modulation active (patient on medication)
- print("\n POST-TREATMENT EVALUATION (with modulation active):")
- neuro_results = {'sparsity': neuro_mgr.get_sparsity() * 100}
- neuro_results['clean'] = evaluate_with_neurosteroid(neuro_model, clean_test_loader, 0.0, 0.0)
- neuro_results['standard'] = evaluate_with_neurosteroid(neuro_model, test_loader, 0.0, 0.0)
- for stress_name, stress_level in CONFIG['extended_stress_levels'].items():
- neuro_results[f'stress_{stress_name}'] = evaluate_with_neurosteroid(neuro_model, test_loader, 0.0, stress_level)
- neuro_results['combined'] = evaluate_with_neurosteroid(neuro_model, test_loader, 1.0, 0.5)
- print(f" Clean: {neuro_results['clean']:.1f}%")
- print(f" Standard: {neuro_results['standard']:.1f}%")
- print(f" Combined stress: {neuro_results['combined']:.1f}%")
- print(f" Extreme stress: {neuro_results['stress_extreme']:.1f}%")
- # Test WITHOUT modulation (medication discontinued)
- print("\n EVALUATION WITHOUT MODULATION (medication discontinued):")
- neuro_model.set_inhibition(1.0, False)
- off_med_combined = evaluate(neuro_model, test_loader, 1.0, 0.5)
- off_med_extreme = evaluate(neuro_model, test_loader, 0.0, 2.5)
- print(f" Combined stress: {off_med_combined:.1f}%")
- print(f" Extreme stress: {off_med_extreme:.1f}%")
- neuro_results['off_medication_combined'] = off_med_combined
- neuro_results['off_medication_extreme'] = off_med_extreme
- # Restore modulation for relapse test
- neuro_model.set_inhibition(CONFIG['neurosteroid_inhibition_strength'],
- CONFIG['neurosteroid_use_tanh'])
- # Relapse simulation
- print("\n RELAPSE SIMULATION (with modulation active):")
- pre_relapse = neuro_results['combined']
- pre_sparsity = neuro_mgr.get_sparsity()
- target_sparsity = min(pre_sparsity + (1 - pre_sparsity) * 0.40, 0.99)
- neuro_mgr.prune_by_magnitude(sparsity=target_sparsity, per_layer=True)
- neuro_mgr.apply_masks()
- post_relapse = evaluate_with_neurosteroid(neuro_model, test_loader, 1.0, 0.5)
- neuro_results['relapse_drop'] = pre_relapse - post_relapse
- print(f" Combined: {pre_relapse:.1f}% → {post_relapse:.1f}% (drop: {neuro_results['relapse_drop']:.1f}%)")
- results['neurosteroid'] = neuro_results
- # ========================================================================
- # COMPREHENSIVE COMPARISON
- # ========================================================================
- print("\n" + "="*80)
- print(" COMPREHENSIVE COMPARISON: ALL TREATMENTS")
- print("="*80)
- treatments = ['untreated', 'ketamine', 'ssri', 'neurosteroid']
- labels = {'untreated': 'Untreated (pruned)', 'ketamine': 'Ketamine-like',
- 'ssri': 'SSRI-like', 'neurosteroid': 'Neurosteroid-like'}
- print(f"\n {'Treatment':<22} {'Sparsity':>10} {'Clean':>8} {'Standard':>10} "
- f"{'Combined':>10} {'Extreme':>10} {'Relapse':>10}")
- print(" " + "-"*85)
- for t in treatments:
- r = results[t]
- relapse = r.get('relapse_drop', 'N/A')
- relapse_str = f"{relapse:.1f}%" if isinstance(relapse, float) else relapse
- print(f" {labels[t]:<22} {r['sparsity']:>9.1f}% {r['clean']:>7.1f}% "
- f"{r['standard']:>9.1f}% {r['combined']:>9.1f}% "
- f"{r['stress_extreme']:>9.1f}% {relapse_str:>10}")
- # Stress resilience profile
- print("\n STRESS RESILIENCE PROFILE:")
- print(f"\n {'Treatment':<22} {'None':>8} {'Moderate':>10} {'High':>8} {'Severe':>8} {'Extreme':>10}")
- print(" " + "-"*70)
- for t in treatments:
- r = results[t]
- print(f" {labels[t]:<22} {r['stress_none']:>7.1f}% {r['stress_moderate']:>9.1f}% "
- f"{r['stress_high']:>7.1f}% {r['stress_severe']:>7.1f}% {r['stress_extreme']:>9.1f}%")
- # ========================================================================
- # ANALYSIS
- # ========================================================================
- print("\n" + "-"*80)
- print(" ANALYSIS")
- print("-"*80)
- ket, ssri, neuro = results['ketamine'], results['ssri'], results['neurosteroid']
- untreated = results['untreated']
- print("\n 1. IMPROVEMENT FROM UNTREATED STATE (Combined Stress):")
- print(f" Ketamine: {untreated['combined']:.1f}% → {ket['combined']:.1f}% (+{ket['combined'] - untreated['combined']:.1f}%)")
- print(f" SSRI: {untreated['combined']:.1f}% → {ssri['combined']:.1f}% (+{ssri['combined'] - untreated['combined']:.1f}%)")
- print(f" Neurosteroid:{untreated['combined']:.1f}% → {neuro['combined']:.1f}% (+{neuro['combined'] - untreated['combined']:.1f}%)")
- print("\n 2. STRUCTURAL VS FUNCTIONAL CHANGES:")
- print(f" Ketamine sparsity: {ket['sparsity']:.1f}% (REDUCED from 95%)")
- print(f" SSRI sparsity: {ssri['sparsity']:.1f}% (UNCHANGED)")
- print(f" Neurosteroid sparsity:{neuro['sparsity']:.1f}% (UNCHANGED)")
- print("\n → Ketamine is the ONLY treatment that adds new connections")
- print("\n 3. EXTREME STRESS RESILIENCE (σ=2.5):")
- print(f" Ketamine: {ket['stress_extreme']:.1f}%")
- print(f" SSRI: {ssri['stress_extreme']:.1f}%")
- print(f" Neurosteroid:{neuro['stress_extreme']:.1f}%")
- print("\n 4. NEUROSTEROID MEDICATION DEPENDENCE:")
- print(f" Combined ON medication: {neuro['combined']:.1f}%")
- print(f" Combined OFF medication: {neuro['off_medication_combined']:.1f}%")
- print(f" Extreme ON medication: {neuro['stress_extreme']:.1f}%")
- print(f" Extreme OFF medication: {neuro['off_medication_extreme']:.1f}%")
- print("\n 5. RELAPSE VULNERABILITY:")
- print(f" Ketamine: {ket['relapse_drop']:.1f}% drop")
- print(f" SSRI: {ssri['relapse_drop']:.1f}% drop")
- print(f" Neurosteroid:{neuro['relapse_drop']:.1f}% drop")
- # ========================================================================
- # CLINICAL INTERPRETATION
- # ========================================================================
- print("\n" + "-"*80)
- print(" CLINICAL INTERPRETATION")
- print("-"*80)
- print("""
- KEY FINDINGS:
- 1. MECHANISM MATTERS: Different antidepressants work through distinct routes
- - Ketamine REBUILDS: Adds new synapses, restores structural density
- - SSRIs REFINE: Strengthen existing pathways via gradual adaptation
- - Neurosteroids STABILIZE: Damp hyperexcitability, bound activity range
- 2. SPEED-DURABILITY TRADEOFF:
- - Ketamine: Fast onset, durable changes (new structure persists)
- - Neurosteroid: Fast onset, medication-dependent (dynamics reset if stopped)
- - SSRI: Slow onset, moderate durability (refined weights persist)
- 3. TREATMENT SELECTION IMPLICATIONS:
- ┌──────────────────────────────────────┬──────────────────────────────────┐
- │ Clinical Scenario │ Suggested Mechanism │
- ├──────────────────────────────────────┼──────────────────────────────────┤
- │ Severe, treatment-resistant MDD │ Ketamine (structural repair) │
- │ Postpartum depression, acute crisis │ Neurosteroid (rapid stabilize) │
- │ Mild-moderate, first-line │ SSRI (gradual, acceptable) │
- │ Recurrent with high relapse risk │ Ketamine (durable structure) │
- │ Hyperexcitable/anxious component │ Neurosteroid (activity damping) │
- └──────────────────────────────────────┴──────────────────────────────────┘
- 4. COMBINATION THERAPY RATIONALE:
- - Ketamine + SSRI: Structural + neuromodulatory benefits
- - Neurosteroid + SSRI: Rapid stabilization while waiting for SSRI onset
- - Ketamine + Psychotherapy: New synapses + activity-guided consolidation
- """)
- # ========================================================================
- # ISO-DOSE COMPARISON EXPERIMENT
- # ========================================================================
- iso_dose_results = run_iso_dose_comparison_experiment(
- base_state_dict, base_masks, untreated_results
- )
- results['iso_dose'] = iso_dose_results
- return results
- # ============================================================================
- # ENTRY POINT
- # ============================================================================
- if __name__ == "__main__":
- print("\n" + "#"*80)
- print("#" + " "*78 + "#")
- print("#" + " MULTI-MECHANISM ANTIDEPRESSANT COMPARISON ".center(78) + "#")
- print("#" + " Ketamine vs SSRI vs Neurosteroid ".center(78) + "#")
- print("#" + " WITH ISO-DOSE FAIR COMPARISON PIPELINE ".center(78) + "#")
- print("#" + " "*78 + "#")
- print("#"*80)
- results = run_multi_mechanism_experiment()
- print("\n" + "="*80)
- print(" EXPERIMENT COMPLETE")
- print("="*80)
- # ========================================================================
- # FINAL SUMMARY: ISO-DOSE FINDINGS
- # ========================================================================
- print("\n" + "="*80)
- print(" FINAL SUMMARY: ISO-DOSE COMPARISON FINDINGS")
- print("="*80)
- if 'iso_dose' in results and 'summary_stats' in results['iso_dose']:
- stats = results['iso_dose']['summary_stats']
- print("\n DOSE CHARACTERISTICS BY TREATMENT:")
- print(f"\n {'Treatment':<15} {'Dose Range':>25} {'Mean Dose':>15}")
- print(" " + "-" * 60)
- for treatment in ['ketamine', 'ssri', 'neurosteroid']:
- s = stats[treatment]
- dose_range = f"{s['dose_range'][0]:.6f} - {s['dose_range'][1]:.6f}"
- print(f" {treatment.capitalize():<15} {dose_range:>25} {s['dose_mean']:>15.6f}")
- print("\n PERFORMANCE RANGE BY TREATMENT:")
- print(f"\n {'Treatment':<15} {'Best Improve':>15} {'Best Relapse':>15} {'Max Efficiency':>15}")
- print(" " + "-" * 65)
- for treatment in ['ketamine', 'ssri', 'neurosteroid']:
- s = stats[treatment]
- print(f" {treatment.capitalize():<15} {s['best_improvement']:>+14.1f}% "
- f"{s['best_relapse']:>14.1f}% {s['max_efficiency']:>15.2f}")
- if 'iso_dose' in results and 'efficiency_analysis' in results['iso_dose']:
- eff = results['iso_dose']['efficiency_analysis']
- print("\n EFFICIENCY COMPARISON:")
- print(f"\n {'Treatment':<15} {'Max Efficiency':>15} {'Mean Efficiency':>16} {'Best Config':>15}")
- print(" " + "-" * 65)
- for treatment in ['ketamine', 'ssri', 'neurosteroid']:
- e = eff[treatment]
- config_str = str(e['best_config']) if e['best_config'] is not None else 'N/A'
- print(f" {treatment.capitalize():<15} {e['max_efficiency']:>15.2f} "
- f"{e['mean_efficiency']:>16.2f} {config_str:>15}")
- print("\n" + "="*80 + "\n")
- # %% [markdown]
- # # The End
068A_antidep_modes_v2.ipynb at commit 7c25bce, no license · at the source
Overview
- Psychiatry, Cheung Ngo Medical Limited, Hong Kong, HKG
Abstract
Background: Major depressive disorder (MDD) is increasingly viewed as a disorder of impaired neural plasticity, yet the mechanisms underlying diverse antidepressant classes - glutamatergic (e.g., ketamine), monoaminergic (e.g., selective serotonin reuptake inhibitors (SSRIs)), and GABAergic (e.g., neurosteroids) - remain incompletely integrated. The objective of this study was to extend a pruning-plasticity model of depression and directly compare, from an identical severely pruned baseline state, the efficacy, stress resilience, durability, and relapse vulnerability of three mechanistically distinct interventions: ketamine-like targeted synaptogenesis, SSRI-like gradual refinement of existing connectivity, and neurosteroid-like tonic inhibition. Computational models offer a controlled means to compare these pathways, but prior work has typically examined single mechanisms.
Methods: We extended a pruning-plasticity model of depression by applying 95% magnitude-based synaptic elimination to overparameterized feed-forward networks trained on a four-class Gaussian classification task. From identical pruned states, three interventions were tested: ketamine-like gradient-guided regrowth (50% reinstatement) with consolidation; SSRI-like prolonged low-learning-rate training with gradual internal noise reduction; and neurosteroid-like global tonic inhibition (30% damping plus tanh activations) with brief consolidation. Outcomes included baseline accuracy, resilience to graded internal activation noise (up to σ = 2.5) plus input perturbation, and relapse vulnerability after an additional 40% pruning.
Results: All treatments restored near-ceiling performance on unchallenged inputs. Ketamine-like synaptogenesis uniquely reduced sparsity (to ~47%) and conferred superior stress resilience (extreme noise accuracy 84.5%) with near-zero relapse drop (−0.2%). SSRI-like refinement improved combined stress accuracy to 83.5% but showed limited extreme noise tolerance (44.0%) and substantial relapse vulnerability (10.8% drop). Neurosteroid-like inhibition achieved rapid combined stress recovery (97.5%) while active, but was state-dependent (decline upon removal) with poor extreme noise buffering (42.5%) and moderate relapse drop (4.1%).
Conclusions: These simulations demonstrate that antidepressants operate through mechanistically distinct routes-structural rebuilding (ketamine), gradual optimization of existing connectivity (SSRIs), or reversible dynamic stabilization (neurosteroids)-yielding
Reproduced under the paper's license (CC BY), from the paper cited above.
Repository
Its files are read in the Code ↔ Paper reader above, with 10 matches between paragraphs and lines of code.
cheungngo/Divergent-Mechanisms-of-Antidepressant-Efficacy
7c25bced1dd4abb138cb9117a07dd0c03211e6d8, 29 January 2026Availability: 1 check, the latest on 30 September 2026: the link answers
- 30 September 2026: the link answers
3 files
- 068A_antidep_modes_v2.ip
ynb , Jupyter, 1,843 lines, 5 matches - 068_antidep_modes.ipynb, Jupyter, 924 lines, 5 matches
- README.md, Text, 622 lines
Tracing map
Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.
What the map holds:
- 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 2 scripts, each with its path and the digest of its content;
- 10 matches between paragraphs of the paper and lines of the code (method lexical-v1);
- neither the text of the paper nor the code itself.
Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.
Data
No dataset and no data link were found in the paper.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 30 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 1 author, 5 keywords, 14 references.
Cite
This paper
Cheung, N. (2026). Divergent Mechanisms of Antidepressant Efficacy: A Unified Computational Comparison of Synaptogenesis, Stabilization, and Tonic Inhibition in a Model of Depression. Cureus, 18(3), e105040. https://
BibTeX
@article{cheung2026diver
author = {Cheung, Ngo},
title = {{Divergent Mechanisms of Antidepressant Efficacy: A Unified Computational Comparison of Synaptogenesis, Stabilization, and Tonic Inhibition in a Model of Depression}},
journal = {Cureus},
year = {2026},
month = mar,
volume = {18},
number = {3},
pages = {e105040},
publisher = {Cureus Inc.},
issn = {2168-8184},
doi = {10.7759/
url = {https://
pmid = {41822249},
pmcid = {PMC12978026}
}
RIS
TY - JOUR
AU - Cheung, Ngo
TI - Divergent Mechanisms of Antidepressant Efficacy: A Unified Computational Comparison of Synaptogenesis, Stabilization, and Tonic Inhibition in a Model of Depression
T2 - Cureus
J2 - Cureus
PY - 2026
DA - 2026/
VL - 18
IS - 3
SP - e105040
SN - 2168-8184
PB - Cureus Inc.
DO - 10.7759/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.7759/
"type": "article-journal",
"title": "Divergent Mechanisms of Antidepressant Efficacy: A Unified Computational Comparison of Synaptogenesis, Stabilization, and Tonic Inhibition in a Model of Depression",
"container-title": "Cureus",
"author": [
{
"family": "Cheung",
"given": "Ngo"
}
],
"container-title-short":
"volume": "18",
"issue": "3",
"page": "e105040",
"DOI": "10.7759/
"PMID": "41822249",
"PMCID": "PMC12978026",
"ISSN": "2168-8184",
"publisher": "Cureus Inc.",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
3,
11
]
]
}
}
The tracing map gets a citation of its own once an author has validated it and it has a DOI.
Similar papers
The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.
- [1] doi:10.3389/fpsyt.2026.1855963 [code]
- Transient entanglement in minimal open XXZ spin chains: a toy-model analogy for microtubule-inspired quantum biology.Journal: Frontiers in psychiatryIn common: clinical / translational, 1 reference, author Ngo Cheung
- [2] doi:10.3389/fpsyt.2026.1752101 [code]
- Case Report: Oral glutamatergic augmentation for trauma-related disorders with fluoxetine-/
bupropion-potentiated dextromethorphan ± piracetam: a four-patient case series. Journal: Frontiers in psychiatryIn common: clinical / translational, author Ngo Cheung - [3] doi:10.1038/s41467-026-77248-y [code]
- The mPFC-reuniens-hippocampu
s pathway links brain circuitry and neural plasticity in antidepressant response. Journal: Nature communicationsIn common: NumPy, depression, 2 references - [4] doi:10.1016/j.isci.2026.116825 [code]
- Social hierarchy shapes behavioral and transcriptional responses to chronic stress and ketamine in male mice.Journal: iScienceIn common: PyTorch, NumPy, depression, 1 reference
- [5] doi:10.1038/s41467-026-75697-z
- Delayed astrocyte development impairs Sema6a-Plxna2/
4-mediated astrocyte-neuron crosstalk and causes depressive-like behavior. Journal: Nature communicationsIn common: depression, 2 references - [6] doi:10.1038/s41398-026-04137-9 [code]
- An integrative mendelian randomisation and drug mechanism framework for target prioritisation and therapeutic repurposing in major depression.Journal: Translational psychiatryIn common: NumPy, depression, clinical / translational, 1 reference
- [7] doi:10.3389/fncom.2026.1799705
- Structural synaptogenesis superior to functional modulation in a pruning-based recurrent network model of OCD.Journal: Frontiers in computational neuroscienceIn common: 2 references
- [8] doi:10.1038/s42003-026-10094-2 [code]
- E/
I imbalance and internal noise cause weak neural representations and face recognition challenges in ASD. Journal: Communications biologyIn common: PyTorch, NumPy, 1 reference - [9] doi:10.1002/hbm.70628 [code]
- EEG Biomarkers for Affective Disorders Diagnosis: An Evaluation and Validation Study.Journal: Human brain mappingIn common: PyTorch, NumPy, depression, clinical / translational
- [10] doi:10.1038/s41380-026-03691-4 [code]
- Breaking the norm: population-scale deviations of brain structure in depression and anxiety.Journal: Molecular psychiatryIn common: PyTorch, NumPy, depression, clinical / translational
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 2 scripts, and 10 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:92ade3c8c5bdf5f4…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
