"""
NeuroSparkSNT -- Symplectic Neural Topology
================================================
C. elegans neural dynamics simulator (NeuroSparkSNT, single model)

Parameter updates:
  gf:         0.08 → 0.05    (PR↓, β→1.0)
  inh:        0.60 → 0.45    (r→-0.42)
  k_hebb:     0.05 → 0.15    (PR↓, PC3sum↑)
  HYSTERESIS: 0.15 → 0.08    (KL↓)
  FWD dwell:  15s  → 20s     (observed ~10.5s → ~14s compensation)

Bug fixes:
  - Micro-jitter 1e-5        (B5 Spearman nan)
  - CommandLayer phase-bias  (show FORWARD/REVERSE instead of PAUSE)
  - Spontaneous REV %10      (KL↓, P(F→R)=0.10 biological)
  - B7 check: delta<0 → ✓   (correct biological criterion)

Additional improvements:
  - γ_t: steer_l also active during FWD (head oscillation / weathervane)
  - sensory_co2 → included in repel_raw (BAG neurons trigger avoidance)
  - state_seq max_dwell_s cap: prevents 119s bug for WC/OW (in bench)
"""

import torch, torch.nn as nn, torch.nn.functional as F
import numpy as np, csv, io, os, glob, warnings, requests
from scipy.stats import pearsonr
from scipy.integrate import solve_ivp
from sklearn.decomposition import PCA
from scipy.signal import detrend as scipy_detrend
import h5py
warnings.filterwarnings("ignore")

# ══════════════════════════════════════════════════════════════════════
# SABITLER
# ══════════════════════════════════════════════════════════════════════
GABA_NEURONS = {
    "DVB","AVL","RIS","RMED","RMEL","RMER","RMEV",
    "DD01","DD02","DD03","DD04","DD05","DD06",
    "VD01","VD02","VD03","VD04","VD05","VD06","VD07",
    "VD08","VD09","VD10","VD11","VD12","VD13",
    "DD1","DD2","DD3","DD4","DD5","DD6",
    "VD1","VD2","VD3","VD4","VD5","VD6","VD7",
    "VD8","VD9","VD10","VD11","VD12","VD13",
}

CIRCUIT_GROUPS = {
    # INPUT (Spark) ─────────────────────────────────────────────────────
    "sensory_attract":    ["AWAL","AWAR","AWCL","AWCR","ASEL","ASER","AIYL","AIYR"],
    "sensory_repel":      ["ASHL","ASHR","ADLL","ADLR","AWBL","AWBR"],
    "sensory_touch_ant":  ["ALML","ALMR","AVM"],
    "sensory_touch_post": ["PLML","PLMR","PVM"],
    "sensory_food":       ["CEPDL","CEPDR","ADEL","ADER","PDEL","PDER"],
    "sensory_temp":       ["AFDL","AFDR"],
    "sensory_oxygen":     ["URXL","URXR","AQR"],
    "sensory_co2":        ["BAGL","BAGR"],          # CO₂↑ → avoidance (like repel)
    # SYSTEM ──────────────────────────────────────────────────────────
    "circuit_forward":    ["AVBL","AVBR","PVCL","PVCR","RIBL","RIBR"],
    "circuit_reverse":    ["AVAL","AVAR","AVDL","AVDR"],   # AIB/RIM in circuit_turn only
    "circuit_turn":       ["AIBL","AIBR","RIML","RIMR","AIZL","AIZR"],
    "circuit_steer":      ["RIAL","RIAR","SMDDL","SMDDR","SMDVL","SMDVR",
                           "RMDDL","RMDDR","RMDL","RMDR"],
    "circuit_pause":      ["RIS","RIPL","RIPR"],
    "circuit_learn":      ["AIYL","AIYR","AIAL","AIAR"],
    # BEHAVIOR ────────────────────────────────────────────────────────
    "motor_forward":      ["DB01","DB02","DB03","DB04","DB05","DB06","DB07",
                           "VB01","VB02","VB03","VB04","VB05","VB06","VB07",
                           "VB08","VB09","VB10","VB11"],
    "motor_reverse":      ["DA01","DA02","DA03","DA04","DA05","DA06","DA07",
                           "DA08","DA09",
                           "VA01","VA02","VA03","VA04","VA05","VA06","VA07",
                           "VA08","VA09","VA10","VA11","VA12"],
    "motor_turn":         ["SMDDL","SMDDR","SMDVL","SMDVR",
                           "RMDDL","RMDDR","RMDL","RMDR"],
    # PHARYNX -- yeme motoru (Cook 2019, Avery 1993)
    "pharynx":            ["M1","M2L","M2R","M3L","M3R","M4","M5",
                           "I1L","I1R","I2L","I2R"],
}

BEHAVIOR_CIRCUITS = {
    "FORWARD":    ["circuit_forward",  "motor_forward"],
    "REVERSE":    ["circuit_reverse",  "motor_reverse"],
    "TURN":       ["circuit_turn",     "motor_turn"],
    "PAUSE":      ["circuit_pause"],
    "CHEMOTAXIS": ["sensory_attract",  "circuit_forward"],
}

PHASE_ACTIVE_FWD = 0
PHASE_RESET      = 1
PHASE_NEUTRAL    = 2
PHASE_ACTIVE_REV  = 3
PHASE_ACTIVE_TURN = 4   # omega turn (Pierce-Shimomura 1999)
PHASE_NAMES      = {0:"FWD", 1:"RESET", 2:"NEUTRAL", 3:"REV", 4:"TURN"}

OW_NEURONS = [
    "ADAL","ADAR","ADEL","ADER","ADFL","ADFR","ADLL","ADLR","AFDL","AFDR",
    "AIAL","AIAR","AIBL","AIBR","AIML","AIMR","AINL","AINR","AIYL","AIYR",
    "AIZL","AIZR","ALA","ALML","ALMR","ALNL","ALNR","AQR",
    "AS1","AS2","AS3","AS4","AS5","AS6","AS7","AS8","AS9","AS10","AS11",
    "ASEL","ASER","ASGL","ASGR","ASHL","ASHR","ASIL","ASIR","ASJL","ASJR",
    "ASKL","ASKR","AUAL","AUAR","AVAL","AVAR","AVBL","AVBR","AVDL","AVDR",
    "AVEL","AVER","AVFL","AVFR","AVG","AVHL","AVHR","AVJL","AVJR","AVKL",
    "AVKR","AVL","AVM","AWAL","AWAR","AWBL","AWBR","AWCL","AWCR",
    "BAGL","BAGR","BDUL","BDUR","CEPDL","CEPDR","CEPVL","CEPVR",
    "DA1","DA2","DA3","DA4","DA5","DA6","DA7","DA8","DA9",
    "DB1","DB2","DB3","DB4","DB5","DB6","DB7",
    "DD1","DD2","DD3","DD4","DD5","DD6",
    "DVA","DVB","DVC","FLPL","FLPR","HSNL","HSNR",
    "I1L","I1R","I2L","I2R","I3","I4","I5","I6",
    "IL1DL","IL1DR","IL1L","IL1R","IL1VL","IL1VR",
    "IL2DL","IL2DR","IL2L","IL2R","IL2VL","IL2VR",
    "LUAL","LUAR","M1","M2L","M2R","M3L","M3R","M4","M5","MCL","MCR","MI",
    "NSML","NSMR","OLLL","OLLR","OLQDL","OLQDR","OLQVL","OLQVR",
    "PDA","PDB","PDEL","PDER","PHAL","PHAR","PHBL","PHBR","PHCL","PHCR",
    "PLML","PLMR","PLNL","PLNR","PQR","PVCL","PVCR","PVDL","PVDR","PVM",
    "PVNL","PVNR","PVPL","PVPR","PVQL","PVQR","PVR","PVT","PVWL","PVWR",
    "RIAL","RIAR","RIBL","RIBR","RICL","RICR","RID","RIFL","RIFR",
    "RIGL","RIGR","RIH","RIML","RIMR","RIPL","RIPR","RIR","RIS",
    "RIVL","RIVR","RMDDL","RMDDR","RMDL","RMDR","RMDVL","RMDVR",
    "RMED","RMEL","RMER","RMEV","RMFL","RMFR","RMGL","RMGR","RMHL","RMHR",
    "SAADL","SAADR","SAAVL","SAAVR","SABD","SABVL","SABVR",
    "SDQL","SDQR","SIADL","SIADR","SIAVL","SIAVR","SIBDL","SIBDR",
    "SIBVL","SIBVR","SMBDL","SMBDR","SMBVL","SMBVR","SMDDL","SMDDR",
    "SMDVL","SMDVR","URADL","URADR","URAVL","URAVR","URBL","URBR",
    "URXL","URXR","URYDL","URYDR","URYVL","URYVR",
    "VA1","VA2","VA3","VA4","VA5","VA6","VA7","VA8","VA9","VA10","VA11","VA12",
    "VB1","VB2","VB3","VB4","VB5","VB6","VB7","VB8","VB9","VB10","VB11",
    "VC1","VC2","VC3","VC4","VC5",
    "VD1","VD2","VD3","VD4","VD5","VD6","VD7","VD8","VD9","VD10","VD11",
    "VD12","VD13",
]

# ══════════════════════════════════════════════════════════════════════
# NEURON REGISTRY
# ══════════════════════════════════════════════════════════════════════
class NeuronRegistry:
    def __init__(self, neuron_names):
        self.names = neuron_names
        self.name2idx = {n:i for i,n in enumerate(neuron_names)}
        self.groups = {
            grp: [self.name2idx[n] for n in ns if n in self.name2idx]
            for grp, ns in CIRCUIT_GROUPS.items()
        }
        self.circuit_idx = sorted(set(
            i for v in self.groups.values() for i in v))
        self.ava = self.groups.get("circuit_reverse",[])[:2]
        self.avb = self.groups.get("circuit_forward",[])[:2]

    def get_circuit_W(self, W_full):
        idx = self.circuit_idx
        return W_full[np.ix_(idx,idx)]

    def local_idx(self, group):
        cm = {v:i for i,v in enumerate(self.circuit_idx)}
        return [cm[i] for i in self.groups.get(group,[]) if i in cm]

    def report(self):
        print(f"\n{'='*55}")
        print("NeuronRegistry -- Cook 2019")
        print(f"{'='*55}")
        for grp, idxs in self.groups.items():
            print(f"  {grp:<28}: {len(idxs):>3}/{len(CIRCUIT_GROUPS[grp])}")
        print(f"  {'Total':<28}: {len(self.circuit_idx)}")
        print(f"  AVA={self.ava}  AVB={self.avb}")
        print(f"{'='*55}")

# ══════════════════════════════════════════════════════════════════════
# COOK 2019
# ══════════════════════════════════════════════════════════════════════
def load_cook2019(path):
    import openpyxl
    wb = openpyxl.load_workbook(path, data_only=True)
    ws = wb["hermaphrodite chemical"]
    row_names = [ws.cell(row=r,column=3).value for r in range(4,304)]
    col_names = [c for c in [ws.cell(row=3,column=c).value
                  for c in range(4,ws.max_column+1)] if c]
    nc_idx = [i for i,c in enumerate(col_names) if c in row_names]
    c2r = {ci:row_names.index(col_names[ci]) for ci in nc_idx}
    W = np.zeros((300,300))
    for ri in range(300):
        for ci in nc_idx:
            v = ws.cell(row=4+ri,column=4+ci).value
            if v not in (None,"."," ","·"):
                try: W[ri,c2r[ci]] = float(v)
                except: pass
    si = sorted(range(300), key=lambda i: row_names[i])
    names = [row_names[i] for i in si]
    W = W[np.ix_(si,si)]
    for i,n in enumerate(names):
        if n in GABA_NEURONS: W[i,:] *= -1
    np.fill_diagonal(W,0)
    sr = np.abs(np.linalg.eigvals(W)).max()
    if sr>1e-8: W *= 1.05/sr
    print(f"  Cook 2019: N=300, {int((W!=0).sum())} synapses")
    return W, names

def add_reciprocal_inhibition(W, ava_idx, avb_idx, strength=0.8):
    if not ava_idx or not avb_idx: return W
    W = W.copy()
    for ai in ava_idx:
        for bi in avb_idx:
            W[bi,ai] -= strength
            W[ai,bi] -= strength*0.9
    sr = np.abs(np.linalg.eigvals(W)).max()
    if sr>1e-8: W *= 1.05/sr
    print(f"  AVA-AVB inhibition added (strength={strength})")
    return W


# ══════════════════════════════════════════════════════════════════════
# STIMULUS AWARENESS LAYER
# ══════════════════════════════════════════════════════════════════════
class StimulusAwarenessLayer:
    """
    Scenario-based sensory → phase bias mapping.

    Translates stimulus type into a directional phase bias injected
    into PhaseController at the state-transition level.

    Biological rationale: each mapping corresponds to a documented
    fast sensory pathway from receptor neurons to command interneurons:
      touch_ant  → ASH → AVA  (Chalfie et al. 1985)
      co2        → BAG → avoidance circuit  (Bretscher 2011)
      oxygen_high → URX → avoidance  (Gray et al. 2004)
      temp_high  → AFD → avoidance  (Glauser et al. 2008)
      temperature → AFD → AIY → klinokinesis TURN  (Iino & Yoshida 2009)
      oxygen_low → BAG → approach  (Zimmer et al. 2009)
      touch_post → PVD → AVB  (Chalfie et al. 1985)

    This layer operates at the PhaseController level (state transitions),
    not at the neural activity level (dx). It is complementary to γ_n
    which biases within-phase dynamics. No conflicts arise because
    the two levels are architecturally separate.

    can_interrupt=True: stimulus can cut active FWD dwell short,
    mimicking fast nociceptive withdrawal reflex pathways.
    """

    BIAS = {
        # stype: (target_phase, strength, can_interrupt_fwd)
        "touch_anterior":  ("REV",  2.5, True),   # ASH nociceptor: hard reversal
        "co2":             ("REV",  1.5, False),   # BAG: graded avoidance
        "oxygen_high":     ("REV",  2.0, True),    # URX: acute avoidance
        "temp_high":       ("REV",  2.0, False),   # AFD: thermal avoidance (Glauser 2008)
        "temperature":     ("TURN", 1.5, False),   # AFD→AIY: klinokinesis reorientation
        "oxygen_low":      ("FWD",  1.0, False),   # BAG: approach
        "touch_posterior": ("FWD",  1.5, False),   # PVD: forward acceleration
    }

    def __init__(self):
        self._stype  = None
        self.target  = None    # "FWD" / "REV" / "TURN" / None
        self.strength = 0.0
        self.can_interrupt = False

    def update(self, stype: str):
        """Call once per step with current effective stimulus type."""
        self._stype = stype
        if stype in self.BIAS:
            t, s, ci = self.BIAS[stype]
            self.target, self.strength, self.can_interrupt = t, s, ci
        else:
            self.target = None
            self.strength = 0.0
            self.can_interrupt = False

    def rev_prob_override(self, base_prob: float) -> float:
        """Adjusted P(→REV) given active stimulus (NEUTRAL decision)."""
        if self.target == "REV":
            return min(0.97, base_prob + self.strength * 0.25)
        elif self.target == "FWD":
            return max(0.02, base_prob - self.strength * 0.15)
        return base_prob

    def wants_turn_from_neutral(self) -> bool:
        """True if awareness directly requests TURN from NEUTRAL."""
        return self.target == "TURN" and self.strength >= 1.0

# ══════════════════════════════════════════════════════════════════════
# PHASE CONTROLLER
# ══════════════════════════════════════════════════════════════════════

# ══════════════════════════════════════════════════════════════════════
# CHAIN TASK SYSTEM -- Goal-Directed Behavior Architecture
# ══════════════════════════════════════════════════════════════════════
# Parallel pipeline: core dynamics intact, modulation added on top
# Each class independent → can be disabled in ablation tests

# Task stage constants
TASK_IDLE     = 0   # Normal pipeline, chain inactive
TASK_APPROACH = 1   # Approach food (LiminalOp + γ_DA)
TASK_CONTACT  = 2   # Food found (transition step)
TASK_CONSUME  = 3   # Ye (pharynx + γ_5HT)
TASK_RESET    = 4   # Cleanup, return to IDLE
TASK_NAMES = {0:"IDLE",1:"APPROACH",2:"CONTACT",3:"CONSUME",4:"RESET"}


class TaskDispatcher:
    """
    Input analizi: reflexive mi, goal-directed mi?
    Goal-directed → ChainTaskPipeline aktive edilir.
    Reflexive → direct base pipeline.
    """
    GOAL_DIRECTED = {"yemek_ye", "yemek_bul"}

    def is_goal_directed(self, stype):
        return stype in self.GOAL_DIRECTED

    def initial_attract_strength(self, stype):
        """Initial odor signal strength for command → distance estimate."""
        return {"yemek_ye": 0.25, "yemek_bul": 0.20}.get(stype, 0.25)

    def has_consume_stage(self, stype):
        """yemek_bul: sadece git+bul. yemek_ye: git+bul+ye."""
        return stype == "yemek_ye"


class LiminalOperator:
    """
    γ_lim -- Spatial Boundary Operator

    Simulates olfactory gradient following with directional feedback:
      FWD  → odor signal increases (approaching food)
      REV  → odor signal decreases (moving away, 50% penalty)
      TURN → signal unchanged (corrective, not navigational)

    Contact is purely signal-based: when sim_attract reaches CONTACT_THR
    (~1.0), the animal is at the food source. No arbitrary step budget needed.

    This directly implements temporal gradient following (Pierce-Shimomura 1999,
    Bargmann & Horvitz 1991): the animal compares current vs past signal
    implicitly through the directional update rule.

    Feedback mechanism:
      - Animal goes the wrong way (REV) → signal drops → less attractive force
      - Omega turn (TURN) → signal holds → ready to try new FWD direction
      - Correct FWD → signal climbs → positive reinforcement
    """
    CONTACT_THR = 0.95   # signal ≈ 1.0 means "at the food source"

    def __init__(self):
        self.initial_attract = 0.0
        self.sim_attract     = 0.0   # current odor signal (distance proxy)
        self.step_size       = 0.0   # signal change per FWD step
        self.progress        = 0.0   # 0→1 normalized

    def initialize(self, initial_attract):
        """
        initial_attract: starting odor strength (0.20–0.40).
        Weaker = farther away. step_size scales so FWD budget
        takes sim_attract from initial → 1.0.
        budget ≈ 300 / (initial + 0.10) FWD steps.
        """
        self.initial_attract = float(initial_attract)
        self.sim_attract     = float(initial_attract)
        budget_steps = max(20, int(300 / (initial_attract + 0.10)))
        self.step_size = (1.0 - initial_attract) / max(1, budget_steps)
        self.progress  = 0.0

    def reset(self):
        self.initial_attract = 0.0
        self.sim_attract     = 0.0
        self.step_size       = 0.0
        self.progress        = 0.0

    def update(self, phase, food_raw_direct):
        """
        Update odor signal based on movement direction.
        This is the feedback loop: wrong direction → signal drops.

        Returns:
          progress    : 0→1 normalized signal progress
          sim_attract : current odor intensity (proxy for proximity)
          contact     : True when sim_attract ≥ CONTACT_THR or food sensor fires
          fwd_boost   : FWD drive magnitude for γ_n
        """
        if phase == PHASE_ACTIVE_FWD:
            # Approaching → signal climbs toward 1.0
            self.sim_attract = min(1.0, self.sim_attract + self.step_size)
        elif phase == PHASE_ACTIVE_REV:
            # Moving away → signal drops (50% of approach rate)
            self.sim_attract = max(self.initial_attract,
                                   self.sim_attract - self.step_size * 0.5)
        # PHASE_ACTIVE_TURN: signal unchanged (corrective movement)

        # Normalized progress: how far from initial to 1.0
        span = max(1e-8, 1.0 - self.initial_attract)
        self.progress = float(np.clip(
            (self.sim_attract - self.initial_attract) / span, 0.0, 1.0))

        # Contact: signal at threshold OR real food sensors fire
        contact = (self.sim_attract >= self.CONTACT_THR) or (food_raw_direct > 0.35)

        # FWD drive: stronger signal → stronger pull toward food
        fwd_boost = self.sim_attract * 0.4

        return self.progress, self.sim_attract, contact, fwd_boost


class BiAmineModulator:
    """
    γ_DA + γ_5HT -- Biamine Antagonism

    Dopamin (DA): hedefe approach sirasinda active
      → increases motor coordination, motivation
    Serotonin (5HT): yemek yendikten sonra active
      → slows locomotion, satiety signal

    Biological dayanak: Sawin et al. 2000,
    Chase & Koelle 2007, Lints & Bhatt 2007
    """
    def __init__(self):
        self.dopamine  = 0.0
        self.serotonin = 0.0

    def reset(self):
        self.dopamine  = 0.0
        self.serotonin = 0.0

    def update(self, task_stage, dt=0.05):
        if task_stage == TASK_APPROACH:
            # Yaklasirken DA yukseliyor, 5HT baskili
            self.dopamine  = min(1.0, self.dopamine  + 1.5 * dt)
            self.serotonin = max(0.0, self.serotonin - 0.3 * dt)
        elif task_stage in [TASK_CONTACT, TASK_CONSUME]:
            # Consuming: DA drops, 5HT rises
            self.dopamine  = max(0.0, self.dopamine  - 1.0 * dt)
            self.serotonin = min(1.0, self.serotonin + 2.5 * dt)
        else:
            # Baseline: both decay
            self.dopamine  = max(0.0, self.dopamine  - 0.4 * dt)
            self.serotonin = max(0.0, self.serotonin - 0.4 * dt)
        return self.dopamine, self.serotonin


class ChainTaskPipeline:
    """
    Goal-directed behavior pipeline.

    Does NOT disrupt v8 core dynamics -- only adds to dx.
    PhaseController continues to run freely at every stage.

    Stage transitions are input-based (not time-based):
      IDLE→APPROACH : goal_directed command detected
      APPROACH→CONTACT: progress>=1 OR food_raw>0.35
      CONTACT→CONSUME : (immediate)
      CONSUME→RESET  : food_raw<0.10 OR serotonin>0.85 (satiety)
      RESET→IDLE     : (immediate, cleanup)
    """
    def __init__(self, registry, dispatcher):
        self.reg        = registry
        self.dispatcher = dispatcher
        self.lim        = LiminalOperator()
        self.amine      = BiAmineModulator()
        self.stage      = TASK_IDLE
        self.has_consume= True     # yemek_ye = True, yemek_bul = False
        # Lokal indeksler
        self.pharynx_l  = registry.local_idx("pharynx")
        self.avb_l      = registry.local_idx("circuit_forward")[:2]
        self.ava_l      = registry.local_idx("circuit_reverse")[:2]
        self.pause_l    = registry.local_idx("circuit_pause")
        self.food_l     = registry.local_idx("sensory_food")
        self.sa_l       = registry.local_idx("sensory_attract")

    def reset(self):
        self.stage = TASK_IDLE
        self.lim.reset()
        self.amine.reset()
        self.has_consume = True

    def get_pool_state(self):
        """Model B icin: pool.forward()'a gecirilecek chain state'i dondur.
        Key 'stage' kullanilir → A ile tutarli (out['chain']['stage'])."""
        return {
            "stage":        self.stage,   # Fix: 'task_stage'→'stage' (A tutarliligi)
            "da_level":     self.amine.dopamine,
            "sht_level":    self.amine.serotonin,
            "sim_attract":  self.lim.sim_attract,  # direct odor signal value
            "has_consume":  self.has_consume,  # yemek_ye=True, yemek_bul=False
        }

    def advance_state(self, phase, food_raw, dt=0.05):
        """Model B icin: sadece chain state'i guncelle, dx'e dokunma.
        Returns: (contact_detected, task_stage)
        Pool operatorler dx modulasyonunu halleder."""
        if self.stage == TASK_IDLE:
            return False, TASK_IDLE

        progress, sim_attract, contact, _ = self.lim.update(phase, food_raw)
        _da, _sht = self.amine.update(self.stage, dt)

        if self.stage == TASK_APPROACH and contact:
            self.stage = TASK_CONTACT   # visible for 1 step (benchmark detects here)
        elif self.stage == TASK_CONTACT:
            # Pool handles pharynx/5HT via TASK_ADD[CONTACT] dispatch
            # but only if has_consume — Pool needs to know this
            # → TASK_ADD is filtered by compute_active_set using task_stage
            # yemek_bul: stage still passes through CONTACT (benchmark records)
            self.stage = TASK_RESET
        elif self.stage == TASK_RESET:
            self.lim.reset()
            self.amine.reset()
            self.stage = TASK_IDLE

        return contact, self.stage

    def activate(self, stype):
        """Komutu al, pipeline'i baslat."""
        init_a = self.dispatcher.initial_attract_strength(stype)
        self.has_consume = self.dispatcher.has_consume_stage(stype)
        self.lim.initialize(init_a)
        self.stage = TASK_APPROACH
        self.amine.reset()

    def step(self, x, dx, phase, food_raw, dt=0.05):
        """
        Mevcut x ve dx'i alir, chain modulasyonunu ekler ve dondurur.
        Hicbir seyi carpmaz -- sadece additive modulasyon.
        """
        if self.stage == TASK_IDLE:
            return dx, {"stage":TASK_IDLE, "DA":0.0, "5HT":0.0,
                        "progress":0.0, "stage_name":"IDLE"}

        progress, sim_attract, contact, fwd_boost = self.lim.update(phase, food_raw)
        da, sht = self.amine.update(self.stage, dt)

        if self.stage == TASK_APPROACH:
            # Drive AVB (FWD motivation via DA)
            if self.avb_l:
                dx[self.avb_l] += da * fwd_boost * (1-x[self.avb_l]) * dt
            # Amplify sensory_attract with simulated odor signal
            if self.sa_l:
                dx[self.sa_l] += sim_attract * 0.3 * (1-x[self.sa_l]) * dt

            if contact:
                # Transition to CONTACT — pharynx fires NEXT step
                # This makes TASK_CONTACT visible in chain_stage_hist
                # so benchmark can detect it with (stages == TASK_CONTACT).any()
                self.stage = TASK_CONTACT

        elif self.stage == TASK_CONTACT:
            # CONTACT is recorded this step (benchmark detects it here)
            # has_consume determines whether to eat (yemek_ye) or just tag (yemek_bul)
            if self.has_consume:
                # yemek_ye: pharynx fires + brief slowing (touch and go)
                if self.pharynx_l:
                    dx[self.pharynx_l] += 1.5 * (1-x[self.pharynx_l]) * dt
                if self.avb_l:
                    dx[self.avb_l] -= 0.3 * x[self.avb_l] * dt
                if self.pause_l:
                    dx[self.pause_l] += 0.5 * (1-x[self.pause_l]) * dt
            # yemek_bul: contact tagged only, no pharynx, no pause
            self.stage = TASK_RESET

        # TASK_CONSUME removed: CONTACT is touch-and-go, directly → RESET
        # One-step pharynx fired, next step RESET

        elif self.stage == TASK_RESET:
            # Cyclic reset: γ_c mechanism takes over
            # Chain equalizes AVA+AVB to push PhaseController toward NEUTRAL
            if self.avb_l and self.ava_l:
                mid = (float(x[self.ava_l].mean()) +
                       float(x[self.avb_l].mean())) / 2.0
                # Pull both toward neutral level
                dx[self.avb_l] += (mid - x[self.avb_l]) * 0.5 * dt
                dx[self.ava_l] += (mid - x[self.ava_l]) * 0.5 * dt
            # Single-step cleanup → IDLE
            self.lim.reset()
            self.amine.reset()
            self.stage = TASK_IDLE

        return dx, {
            "stage":      self.stage,
            "stage_name": TASK_NAMES.get(self.stage,"?"),
            "progress":   progress,
            "DA":         da,
            "5HT":        sht,
            "sim_attract":sim_attract,
        }

class PhaseController:
    """
    Semi-Markov phase kontrolcusu.

     degisiklikleri:
    - HYSTERESIS: 0.15→0.08 (daha sik gecis, KL↓)
    - FWD dwell μ: 15→20 (gozlemlenen kisalmayi telafi)
    - Spontaneous REV: NEUTRAL'da %10 olasilikla REV
      (C. elegans P(F→R)=0.10 biological gereksinimi)
    """
    DWELL = {
        PHASE_ACTIVE_FWD:  (20.0, 0.6),
        PHASE_ACTIVE_REV:  (2.0,  0.7),
        PHASE_RESET:       (None, None),
        PHASE_NEUTRAL:     (0.8,  0.5),
        PHASE_ACTIVE_TURN: (0.5,  0.4),  # omega turn ~0.5s (Pierce-Shimomura 1999)
    }
    TURN_PROB_AFTER_REV = 0.60   # 60% of reversals end with omega turn
    RESET_S         = 1.0
    REFRACTORY_S    = 1.5
    HYSTERESIS      = 0.08   # 0.15→0.08 (KL↓)
    SPONTANEOUS_REV = 0.10   # P(F→R) = 0.10, Kato 2015

    def __init__(self, dt=0.05):
        self.dt          = dt
        self.phase       = PHASE_NEUTRAL
        self.dwell_rem   = 0
        self.ref_to_fwd  = 0
        self.ref_to_rev  = 0
        self.acil_rev    = False
        self._rng        = np.random.default_rng(0)

    def _dwell(self, phase):
        if phase == PHASE_RESET:
            return max(1, int(self.RESET_S/self.dt))
        mu, sig = self.DWELL[phase]
        sec = float(np.clip(self._rng.lognormal(np.log(mu), sig), 0.5, 60.0))
        return max(1, int(sec/self.dt))

    def reset(self):
        self.phase      = PHASE_NEUTRAL
        self.dwell_rem  = self._dwell(PHASE_NEUTRAL)
        self.ref_to_fwd = 0
        self.ref_to_rev = 0
        self.acil_rev   = False

    def _go(self, new_phase):
        old = self.phase
        self.phase     = new_phase
        self.dwell_rem = self._dwell(new_phase)
        # FWD dwell end → prevent immediate REV transition
        # REV dwell end → prevent immediate FWD transition
        if old == PHASE_ACTIVE_FWD:
            self.ref_to_rev = int(self.REFRACTORY_S/self.dt)
        elif old == PHASE_ACTIVE_REV:
            self.ref_to_fwd = int(self.REFRACTORY_S/self.dt)

    def update(self, ava_act, avb_act, repel_act=0.0, touch_ant=0.0,
               awareness=None):
        """
        awareness: StimulusAwarenessLayer or None.
        Injects scenario-based phase bias at the transition level.
        """
        prev = self.phase
        if self.ref_to_fwd > 0: self.ref_to_fwd -= 1
        if self.ref_to_rev > 0: self.ref_to_rev -= 1
        if self.dwell_rem  > 0: self.dwell_rem  -= 1

        # Fast nociceptive interrupt: strong aversive OR awareness override
        # Biologically: ASH/URX → fast avoidance pathway (< 1 dwell step)
        aw_interrupt = (awareness is not None and
                        awareness.can_interrupt and
                        self.ref_to_rev == 0)
        if (repel_act > 0.5 or touch_ant > 0.5 or aw_interrupt) and            self.phase == PHASE_ACTIVE_FWD:
            self.acil_rev = True
            self._go(PHASE_RESET)
            return self.phase, True

        if self.dwell_rem > 0:
            return self.phase, False

        # ── Dwell elapsed — time-based exit ─────────────────────────
        if self.phase == PHASE_ACTIVE_FWD:
            self._go(PHASE_RESET)

        elif self.phase == PHASE_ACTIVE_REV:
            if self._rng.random() < self.TURN_PROB_AFTER_REV:
                self._go(PHASE_ACTIVE_TURN)
            else:
                self._go(PHASE_RESET)

        elif self.phase == PHASE_ACTIVE_TURN:
            self._go(PHASE_RESET)

        elif self.phase == PHASE_RESET:
            self._go(PHASE_NEUTRAL)

        elif self.phase == PHASE_NEUTRAL:
            if self.acil_rev and self.ref_to_rev == 0:
                self.acil_rev = False
                self._go(PHASE_ACTIVE_REV)

            # ── Awareness-driven direct TURN from NEUTRAL ────────────
            # Biologically: klinokinesis — increased turn rate when
            # moving away from thermal/chemical optimum (Iino 2009)
            elif awareness is not None and awareness.wants_turn_from_neutral():
                self._go(PHASE_ACTIVE_TURN)

            elif avb_act > ava_act + self.HYSTERESIS and self.ref_to_fwd == 0:
                self._go(PHASE_ACTIVE_FWD)

            elif ava_act > avb_act + self.HYSTERESIS and self.ref_to_rev == 0:
                self._go(PHASE_ACTIVE_REV)

            else:
                # Awareness biases P(→REV) via scenario mapping
                base_rev_p = self.SPONTANEOUS_REV
                rev_p = (awareness.rev_prob_override(base_rev_p)
                         if awareness is not None else base_rev_p)
                if self._rng.random() < rev_p and self.ref_to_rev == 0:
                    self._go(PHASE_ACTIVE_REV)
                elif self.ref_to_fwd == 0:
                    self._go(PHASE_ACTIVE_FWD)
                else:
                    self.dwell_rem = self._dwell(PHASE_NEUTRAL)

        return self.phase, (prev != self.phase)

# ══════════════════════════════════════════════════════════════════════
# OPERATOR LAYER -- 7 Operator
# ══════════════════════════════════════════════════════════════════════
class OperatorLayer(nn.Module):
    """
    Eight-operator monolith pipeline (Model A).
    All operators run every step; TURN/FWD/REV gated by phase argument.

    INPUT  (Spark): γ_f   exploration noise (gf=0.05)
    INPUT  (Spark): γ_d   stimulus detection (threshold > STIM_THR)
    SYSTEM        : γ_n   decision bias toward FWD/REV (threshold > SENS_THR)
    SYSTEM        : γ_c   cyclic reset (excluded from ACTIVE_FWD/REV/TURN)
    SYSTEM        : γ_t   tonic homeostasis + weathervane head oscillation
    BEHAVIOR      : γ_fwd forward locomotion only (hard lock, inh=0.45)
    BEHAVIOR      : γ_rev reverse locomotion only (hard lock)
    BEHAVIOR      : γ_turn omega turn — PHASE_ACTIVE_TURN only (g_turn=1.20)

    γ_learn: side-effect inside forward() — slowly updates T_c via circuit_learn.
    """
    SENS_THR = 0.25    # γ_n threshold: ignore noise below this level
    STIM_THR = 0.20    # γ_d threshold: filter weak stimuli

    def __init__(self):
        super().__init__()
        self.g_fwd = 1.50; self.g_rev = 1.50; self.g_turn = 1.20
        self.gf    = 0.05   # 0.08→0.05 (PR↓, β→1.0)
        self.gn    = 1.50
        self.gd    = 0.50
        self.gc    = 1.20
        self.gt    = 0.15
        self.tau   = 0.80; self.tauf  = 0.10; self.taud = 5.00
        self.zeta  = 0.15
        self.inh   = 0.45   # 0.60→0.45 (r→-0.42)
        self.k_hebb = 0.15  # 0.05→0.15 (PR↓, PC3sum↑)
        # T_c thermotaxis: yetistirme sicakligi hafizasi (Iino & Yoshida 2009)
        # Normalize edilmis: 20°C ≈ 0.40 (model olceginde)
        self.T_c         = 0.40   # cultivation temp memory (normalized 20°C)
        self.alpha_learn = 0.001  # T_c adaptation rate (slow, minute-scale)

    def _gamma_turn(self, x, turn_l, motor_turn_l, ava_l, avb_l, dt):
        """γ_turn: omega turn operator — PHASE_ACTIVE_TURN only.
        Activates circuit_turn (AIB/RIM) + motor_turn (SMD/RMD).
        No FWD/REV conflict: called only when PHASE_ACTIVE_TURN.
        Reference: Pierce-Shimomura 1999
        """
        d = torch.zeros(len(x))
        if turn_l:       d[turn_l]       += self.g_turn*(1-x[turn_l])*dt
        if motor_turn_l: d[motor_turn_l] += self.g_turn*0.6*(1-x[motor_turn_l])*dt
        if ava_l:        d[ava_l]        -= self.inh*x[ava_l]*dt
        if avb_l:        d[avb_l]        -= self.inh*0.5*x[avb_l]*dt
        return d

    def _gamma_learn(self, learn_l, temp_raw, x, dt):
        """γ_learn: thermotaxis memory via circuit_learn (AIY/AIZ).
        Updates T_c toward current temperature experience (Iino & Yoshida 2009).
        Side-effect only — no dx contribution.
        """
        if not learn_l or float(temp_raw) <= 0.10:
            return
        la = float(x[learn_l].mean())
        self.T_c = float(np.clip(
            self.T_c + self.alpha_learn*(float(temp_raw)-self.T_c)*la*dt,
            0.10, 0.90))

    def forward(self, x, theta, W, H,
                ava_l, avb_l, fwd_m, rev_m,
                steer_l, turn_l, pause_l,
                sa_l, sr_l, tant_l, tpost_l,
                food_l, temp_l, co2_l, o2_l,
                phase, rng, dt=0.05,
                learn_l=None, motor_turn_l=None):
        # OperatorLayer: Model A monolith pipeline
        # (OperatorPool uses pool/dispatch architecture; this does not)
        N = len(x)

        # Raw sensor activations
        attract_raw = float(x[sa_l].mean())    if sa_l    else 0.0
        repel_raw   = float(x[sr_l].mean())    if sr_l    else 0.0
        co2_raw     = float(x[co2_l].mean())   if co2_l   else 0.0
        o2_raw      = float(x[o2_l].mean())    if o2_l    else 0.0
        tant_raw    = float(x[tant_l].mean())  if tant_l  else 0.0
        tpost_raw   = float(x[tpost_l].mean()) if tpost_l else 0.0
        food_raw    = float(x[food_l].mean())  if food_l  else 0.0
        temp_raw    = float(x[temp_l].mean())  if temp_l  else 0.0
        # CO₂ kacinmayi repel'e ekle
        repel_eff_raw = repel_raw + co2_raw * 0.80  # CO₂ full weight → REV

        # ── γ_f: kesif gurultusu ─────────────────────────────────────
        ns = 1.5 if phase==PHASE_RESET else 1.2 if phase==PHASE_NEUTRAL else 1.0
        d_f = self.gf * ns * \
              torch.randn(N, generator=rng)*np.sqrt(dt/self.tauf)

        # ── γ_d: uyaran amplifikasyonu (esikli) ──────────────────────
        d_d = torch.zeros(N)
        if attract_raw > self.STIM_THR and sa_l:
            d_d[sa_l] += self.gd * attract_raw * (1-x[sa_l]) * dt
        if repel_eff_raw > self.STIM_THR and sr_l:
            d_d[sr_l] += self.gd * repel_eff_raw * (1-x[sr_l]) * dt
        if food_raw > self.STIM_THR and sa_l:
            d_d[sa_l] += self.gd * 0.3 * food_raw * (1-x[sa_l]) * dt

        # ── γ_n: karar bias (esikli -- gurultuye tepki vermez) ────────
        d_n = torch.zeros(N)
        a_eff = max(0.0, attract_raw    - self.SENS_THR)
        r_eff = max(0.0, repel_eff_raw  - self.SENS_THR)
        tp_eff= max(0.0, tpost_raw      - self.SENS_THR)
        ta_eff= max(0.0, tant_raw       - self.SENS_THR)
        # O2 avoidance: URX/AQR activated by high O2 (21%) → avoidance
        # Gray 2004, Zimmer 2009: URX rises with O2 increase → aerotaxis
        # oxygen_high stim → URX active → o2_eff → rev_bias
        # oxygen_low stim → BAG active → computed via co2_raw
        o2_eff  = max(0.0, o2_raw - self.SENS_THR) * 0.4

        # Temperature extreme: above HIGH_TEMP_THR → reversal tendency (Glauser 2008)
        # Mild temperature (temp_raw < 0.45) only drives weathervane steering
        # Extreme sicaklik (temp_raw > 0.45) reversal tetikler
        HIGH_TEMP_THR = 0.45
        temp_stress_eff = max(0.0, temp_raw - HIGH_TEMP_THR) * 1.5

        # Thermotaxis T_c memory: temperature error signal (Iino & Yoshida 2009)
        # T > T_c → soguga kac (rev_bias); T < T_c → isiya yonel (fwd_bias)
        # Isothermal tracking (T ≈ T_c): steer_l already active in γ_t
        temp_err = temp_raw - self.T_c
        tc_fwd_eff = max(0.0, -temp_err - 0.05) * 0.35  # soguga kac
        tc_rev_eff = max(0.0,  temp_err - 0.05) * 0.55  # thermal avoidance

        food_bias_eff = max(0.0, food_raw - self.SENS_THR) * 0.30
        fwd_bias = a_eff*0.4 + tp_eff*0.5 + tc_fwd_eff + food_bias_eff
        rev_bias = r_eff*0.5 + ta_eff*0.6 + o2_eff + temp_stress_eff + tc_rev_eff
        if avb_l and fwd_bias > 0:
            d_n[avb_l] = self.gn * fwd_bias * (1-x[avb_l]) * dt
        if ava_l and rev_bias > 0:
            d_n[ava_l] = self.gn * rev_bias * (1-x[ava_l]) * dt

        # ── γ_t: tonik homeostaz + dopamin + bas salinimi ────────────
        dopamin_brake = float(np.clip(food_raw*0.5, 0, 0.35))
        d_t = self.gt * (x.mean()-x) * dt
        # Head oscillation (weathervane): also active during FWD
        # Steer neurons adjust direction based on odor gradient
        if steer_l:
            steer_drive = temp_raw if temp_raw > self.STIM_THR else attract_raw*0.3
            if steer_drive > 0.05:
                d_t[steer_l] += self.gt * steer_drive * (1-x[steer_l]) * dt * 0.6

        # ── γ_c: cyclic reset ─────────────────────────────────────────
        # CLOSED during ACTIVE phases (no conflict with γ_fwd/γ_rev)
        if phase == PHASE_RESET:
            gc_eff = 4.0
        elif phase == PHASE_NEUTRAL:
            gc_eff = self.gc * 0.3
        else:
            gc_eff = 0.0   # ACTIVE_FWD / ACTIVE_REV: kapali
        d_c = -gc_eff * x * (x > theta).float() * dt
        # Omega turn interneurons: active during RESET phase
        if phase == PHASE_RESET and turn_l:
            d_c[turn_l] += self.gc * 0.4 * (1-x[turn_l]) * dt
        # Pause/sleep circuit: active during NEUTRAL phase
        if phase == PHASE_NEUTRAL and pause_l:
            d_c[pause_l] += self.gc * 0.15 * (1-x[pause_l]) * dt

        # ── Decay ─────────────────────────────────────────────────────
        d_decay = -self.gt * (3.0 if phase==PHASE_RESET else 1.0) * x * dt

        # ── Connectome + Hebb koordinasyon ────────────────────────────
        d_tau  = (1/self.tau)*(-x + torch.tanh(W@x))*dt
        # H symmetric W: coordinates active neurons with each other
        # k_hebb=0.15: PR↓, PC3sum↑ (neurons settle onto manifold)
        d_hebb = self.k_hebb * torch.tanh(H@x) * dt

        # ── γ_fwd (Nexter) + γ_rev (Reverser) -- hard lock ────────────
        d_fwd = torch.zeros(N)
        d_rev = torch.zeros(N)

        if phase == PHASE_ACTIVE_FWD:
            exc = self.g_fwd * (1-dopamin_brake)
            if avb_l: d_fwd[avb_l] = exc * (1-x[avb_l]) * dt
            if fwd_m:  d_fwd[fwd_m] = exc * 0.5 * (1-x[fwd_m]) * dt
            # γ_rev CLOSED: suppress active AVA
            if ava_l: d_fwd[ava_l] = -self.inh * x[ava_l] * dt

        elif phase == PHASE_ACTIVE_REV:
            if ava_l: d_rev[ava_l] = self.g_rev * (1-x[ava_l]) * dt
            if rev_m:  d_rev[rev_m] = self.g_rev * 0.5 * (1-x[rev_m]) * dt
            # γ_fwd CLOSED: suppress active AVB
            if avb_l: d_rev[avb_l] = -self.inh * x[avb_l] * dt

        # RESET / NEUTRAL: her iki operator sifir (sistem bosalir)

        dx = (d_tau + d_hebb + d_f + d_d + d_t + d_c +
              d_decay + d_fwd + d_rev + d_n)
                # γ_turn: omega turn (no FWD/REV conflict)
        if phase == PHASE_ACTIVE_TURN:
            dx += self._gamma_turn(x, turn_l, motor_turn_l, ava_l, avb_l, dt)
        # γ_learn: T_c side-effect only
        self._gamma_learn(learn_l, temp_raw, x, dt)
        return dx, {"attract":attract_raw, "repel":repel_eff_raw,
                    "food":food_raw, "dopamin_brake":dopamin_brake,
                    "fwd_bias":fwd_bias, "rev_bias":rev_bias}

# ══════════════════════════════════════════════════════════════════════
# YARDIMCI SINIFLAR
# ══════════════════════════════════════════════════════════════════════
class _StimulusLayer(nn.Module):
    STIM_MAP = {
        "food_odor":      "sensory_attract",
        "noxious":        "sensory_repel",
        "touch_anterior": "sensory_touch_ant",
        "touch_posterior":"sensory_touch_post",
        "food_present":   "sensory_food",
        "temperature":    "sensory_temp",    # hafif sicaklik → weathervane
        "temp_high":      "sensory_temp",    # high temperature → avoidance
        "co2":            "sensory_co2",     # CO₂↑ → kacinma
        "oxygen_high":    "sensory_oxygen",  # O₂↑ 21% → avoidance (URX detects high O2), Gray 2004)
        # "oxygen_low": removed from STIM_MAP — sensory_co2 contributes
        #   wrong direction (rev_bias). FWD drive handled by StimulusAwarenessLayer only. O₂↓ → avoidance (BAG neurons detect low O2, Zimmer 2009)
        "yemek_ye":      None,   # Goal-directed: TaskDispatcher activates
        "yemek_bul":     None,   # Goal-directed: APPROACH+CONTACT only (no consume)
        "none":          None,
    }
    def __init__(self, dim_n, registry):
        super().__init__()
        self.dim_n=dim_n; self.reg=registry

    def forward(self, n, stype, strength=0.6, t=0.0,
                stim_on=5.0, stim_off=15.0, decay=0.20):
        if stype=="none" or not (stim_on<=t<=stim_off):
            return torch.zeros(self.dim_n)
        tgt = self.STIM_MAP.get(stype)
        if tgt is None: return torch.zeros(self.dim_n)
        s = torch.zeros(self.dim_n)
        mag = float(strength*np.exp(-decay*(t-stim_on)))
        cm = {v:i for i,v in enumerate(self.reg.circuit_idx)}
        for gi in self.reg.groups.get(tgt,[]):
            if gi in cm: s[cm[gi]] = mag
        return s


class _CommandLayer(nn.Module):
    """
    Phase-biased CommandLayer ( fix):
    ACTIVE phaselarda phase-uyumlu behaviora +1.5 bonus.
    Onceki versiyonda PAUSE display sorunu vardi.
    """
    def __init__(self, registry):
        super().__init__()
        self.reg   = registry
        self.bc    = BEHAVIOR_CIRCUITS
        self.bnames= list(BEHAVIOR_CIRCUITS.keys())

    def forward(self, n, phase):
        if phase in [PHASE_RESET, PHASE_NEUTRAL]:
            logits = torch.zeros(len(self.bnames))
            if "PAUSE" in self.bnames:
                logits[self.bnames.index("PAUSE")] = 2.0
            if "TURN" in self.bnames:
                logits[self.bnames.index("TURN")] = 0.5
            return F.softmax(logits, dim=0)

        # Active phases: compute logit from neuron activation
        logits = torch.zeros(len(self.bnames))
        for j, bn in enumerate(self.bnames):
            vals = [n[self.reg.local_idx(g)].mean()
                    for g in self.bc.get(bn,[])
                    if self.reg.local_idx(g)]
            if vals: logits[j] = torch.stack(vals).mean()

        # Phase bias: consistent behavior during ACTIVE phase (PAUSE display fix)
        if phase == PHASE_ACTIVE_FWD and "FORWARD" in self.bnames:
            logits[self.bnames.index("FORWARD")] += 1.5
        elif phase == PHASE_ACTIVE_REV and "REVERSE" in self.bnames:
            logits[self.bnames.index("REVERSE")] += 1.5
        elif phase == PHASE_ACTIVE_TURN and "TURN" in self.bnames:
            logits[self.bnames.index("TURN")] += 1.5

        return F.softmax(logits, dim=0)


class _SoftClamp(nn.Module):
    """
    SoftClamp: theta hedefi phasea gore ayarlanir.
    ACTIVE phaselarda HIGH theta (γ_c motor neuronlarina dokunmaz).
    RESET'te LOW theta (agresif temizleme).
    """
    def __init__(self, eta=0.03, mu=0.90, lo=0.28, hi=0.68, init=0.47):
        super().__init__()
        self.eta=eta; self.mu=mu; self.lo=lo; self.hi=hi; self.init=init

    def forward(self, x, theta, vt, phase, attract=0.0, repel=0.0):
        if phase in (PHASE_ACTIVE_FWD, PHASE_ACTIVE_REV, PHASE_ACTIVE_TURN):
            # Active locomotion: high theta → γ_c stays off motor neurons
            td = 0.65
        elif phase == PHASE_RESET:
            td = 0.15   # aggressive cleanup
        else:           # NEUTRAL
            td = float(np.clip(0.40+attract*0.05-repel*0.05, self.lo, self.hi))
        vn = self.mu*vt + (1-self.mu)*(x-td)
        return torch.clamp(theta+self.eta*vn, self.lo, self.hi), vn

# ══════════════════════════════════════════════════════════════════════
# NeuroSparkSNT -- TAM MODEL
# ══════════════════════════════════════════════════════════════════════
class NeuroSparkSNT(nn.Module):
    def __init__(self, W_circuit, registry, dt=0.05):
        super().__init__()
        self.reg   = registry
        self.dim_n = W_circuit.shape[0]
        self.dt    = dt
        self.bnames= list(BEHAVIOR_CIRCUITS.keys())

        W_t = torch.tensor(W_circuit, dtype=torch.float32)
        self.register_buffer("W", W_t)
        self.register_buffer("H", (W_t+W_t.T)/2.0)

        li = registry.local_idx
        self.ava_l  = li("circuit_reverse")[:2]
        self.avb_l  = li("circuit_forward")[:2]
        self.fwd_m  = li("motor_forward")
        self.rev_m  = li("motor_reverse")
        self.steer_l= li("circuit_steer")
        self.turn_l = li("circuit_turn")
        self.pause_l= li("circuit_pause")
        self.sa_l   = li("sensory_attract")
        self.sr_l   = li("sensory_repel")
        self.tant_l = li("sensory_touch_ant")
        self.tpost_l= li("sensory_touch_post")
        self.food_l = li("sensory_food")
        self.temp_l = li("sensory_temp")
        self.co2_l  = li("sensory_co2")
        self.o2_l   = li("sensory_oxygen")

        self.stim  = _StimulusLayer(self.dim_n, registry)
        self.ops   = OperatorLayer()
        self.cmd   = _CommandLayer(registry)
        self.clamp = _SoftClamp()
        self.phase_ctrl = PhaseController(dt=dt)
        # Chain Task Pipeline
        self.dispatcher = TaskDispatcher()
        self.chain      = ChainTaskPipeline(registry, self.dispatcher)
        self.awareness  = StimulusAwarenessLayer()
        self.pharynx_l   = li("pharynx")
        self.learn_l     = li("circuit_learn")
        self.motor_turn_l= li("motor_turn")

        print(f"  NeuroSparkSNT: dim={self.dim_n}")
        print(f"  AVA={self.ava_l}  AVB={self.avb_l}")
        print(f"  FWD_m={len(self.fwd_m)}  REV_m={len(self.rev_m)}  "
              f"Steer={len(self.steer_l)}  Pharynx={len(self.pharynx_l)}")
        print(f"  ChainTask: TaskDispatcher+LiminalOp+BiAmine ✓")

    def _fresh(self, rng):
        self.phase_ctrl._rng = np.random.default_rng(
            int(torch.randint(0,99999,(1,),generator=rng).item()))
        self.phase_ctrl.reset()
        self.chain.reset()
        return {
            "x":     torch.rand(self.dim_n, generator=rng)*0.15+0.10,
            "theta": torch.full((self.dim_n,), self.clamp.init),
            "vt":    torch.zeros(self.dim_n),
            "vx":    torch.zeros(self.dim_n),
        }

    def step(self, state, stype="none", t=0.0, rng=None):
        if rng is None: rng=torch.Generator()
        x,theta,vt,vx = (state["x"],state["theta"],
                          state["vt"],state["vx"])

        ava_a = float(x[self.ava_l].mean())  if self.ava_l  else 0.0
        avb_a = float(x[self.avb_l].mean())  if self.avb_l  else 0.0
        rep_a = float(x[self.sr_l].mean())   if self.sr_l   else 0.0
        tan_a = float(x[self.tant_l].mean()) if self.tant_l else 0.0

        # TaskDispatcher: is this a goal-directed command?
        if self.dispatcher.is_goal_directed(stype) and self.chain.stage == TASK_IDLE:
            self.chain.activate(stype)
        effective_stype = stype if not self.dispatcher.is_goal_directed(stype) else "food_odor"

        # StimulusAwarenessLayer: inject scenario bias into PhaseController
        self.awareness.update(effective_stype)
        phase, _ = self.phase_ctrl.update(ava_a, avb_a, rep_a, tan_a,
                                           awareness=self.awareness)

        stim = self.stim(x, effective_stype, t=t)
        dx, info = self.ops(
            x, theta, self.W, self.H,
            self.ava_l, self.avb_l, self.fwd_m, self.rev_m,
            self.steer_l, self.turn_l, self.pause_l,
            self.sa_l, self.sr_l, self.tant_l, self.tpost_l,
            self.food_l, self.temp_l, self.co2_l, self.o2_l,
            phase, rng, self.dt,
            learn_l=self.learn_l, motor_turn_l=self.motor_turn_l)
        dx = dx + stim*self.dt

        # ── Paralel Chain Pipeline -- v8 dx ustune, vx_new'den ONCE ──────
        # v8 ran fully. Chain adds on top. SoftClamp/PhaseCtrl run independently.
        food_raw_direct = float(x[self.food_l].mean()) if self.food_l else 0.0
        dx, chain_info  = self.chain.step(x, dx, phase, food_raw_direct, self.dt)

        # Momentum + state update (combined v8 + chain dx)
        vx_new = (1-self.ops.zeta)*dx + self.ops.zeta*vx
        x_new  = torch.clamp(x+vx_new, 0.0, 1.0)

        # Micro-jitter: B5 Spearman nan fix
        with torch.no_grad():
            x_new = x_new + torch.randn_like(x_new) * 1e-5
            x_new = torch.clamp(x_new, 0.0, 1.0)

        beh    = self.cmd(x_new, phase)
        th_new, vt_new = self.clamp(
            x_new, theta, vt, phase, info["attract"], info["repel"])

        return {"x":x_new, "theta":th_new, "vt":vt_new, "vx":vx_new,
                "beh":beh, "phase":phase, "info":info, "chain":chain_info}

    def simulate(self, n_steps=1200, stype="food_odor",
                 stim_on=5.0, stim_off=15.0, seed=42, verbose=True):
        rng = torch.Generator(); rng.manual_seed(seed)
        torch.manual_seed(seed)
        t_arr = np.arange(n_steps)*self.dt
        state = self._fresh(rng)
        X   = torch.zeros(n_steps, self.dim_n)
        B   = torch.zeros(n_steps, len(self.bnames))
        ph  = np.zeros(n_steps, dtype=int)
        chain_stage_hist = np.zeros(n_steps, dtype=int)

        if verbose:
            print(f"\n{'='*68}")
            print(f"NeuroSparkSNT  stype={stype}  [{stim_on}-{stim_off}s]")
            print(f"{'='*68}")
            print(f"{'t':>6}|{'Phase':<8}|{'Behavior':<12}|"
                  f"{'AVA':>6}|{'AVB':>6}|{'r(50)':>8}|{'Chain'}")
            print("-"*55)

        for i, t in enumerate(t_arr):
            st = stype if stim_on<=t<=stim_off else "none"
            out = self.step(state, st, t, rng)
            state = out
            X[i]=out["x"]; B[i]=out["beh"]; ph[i]=out["phase"]
            chain_stage_hist[i]=out["chain"]["stage"]

            if verbose and (i%240==0 or i==n_steps-1):
                dom = self.bnames[torch.argmax(out["beh"]).item()]
                pn  = PHASE_NAMES.get(out["phase"],"?")
                aa  = float(out["x"][self.ava_l].mean()) if self.ava_l else 0.0
                ba  = float(out["x"][self.avb_l].mean()) if self.avb_l else 0.0
                if i>=50 and self.ava_l and self.avb_l:
                    a=X[max(0,i-50):i, self.ava_l].mean(1)
                    b=X[max(0,i-50):i, self.avb_l].mean(1)
                    r=(torch.corrcoef(torch.stack([a,b]))[0,1].item()
                       if a.std()>1e-6 and b.std()>1e-6 else float('nan'))
                else: r=float('nan')
                chain_s = TASK_NAMES.get(out["chain"]["stage"],"?")
                print(f"{t:6.1f}|{pn:<8}|{dom:<12}|"
                      f"{aa:6.3f}|{ba:6.3f}|{r:8.4f}|{chain_s}")

        ava_avb_r = float('nan')
        if self.ava_l and self.avb_l:
            a=X[:, self.ava_l].mean(1).detach().numpy()
            b=X[:, self.avb_l].mean(1).detach().numpy()
            if a.std()>1e-6 and b.std()>1e-6:
                ava_avb_r, _ = pearsonr(a, b)

        pp  = {PHASE_NAMES[p]:float((ph==p).mean()*100) for p in range(5)}
        dom_b = self.bnames[B.mean(0).argmax().item()]

        if verbose:
            print(f"\n  Dominant behavior : {dom_b}")
            print(f"  AVA-AVB r       : {ava_avb_r:+.4f}  (ground truth: -0.420)")
            print(f"  Phases          : "+" ".join(
                f"{k}:{v:.0f}%" for k,v in pp.items()))
            print(f"{'='*68}")

        return {"X":X,"B":B,"ava_avb_r":ava_avb_r,
                "dom":dom_b,"phase_pcts":pp,
                "chain_stage_hist": chain_stage_hist, "ph": ph}

# ══════════════════════════════════════════════════════════════════════
# KARSILASTIRMA MODELLERI
# ══════════════════════════════════════════════════════════════════════
class WilsonCowanBaseline:
    def __init__(self, W):
        sr = np.abs(np.linalg.eigvals(W)).max()
        self.W = torch.tensor(W*(1.05/sr) if sr>1e-8 else W,
                               dtype=torch.float32)
        self.N = W.shape[0]

    def run(self, n_steps=1200, seed=42):
        rng = torch.Generator(); rng.manual_seed(seed)
        x = torch.rand(self.N, generator=rng)*0.2
        X = torch.zeros(n_steps, self.N)
        for i in range(n_steps):
            X[i] = x
            x = torch.clamp(
                x+(1/0.8)*(-x+1/(1+torch.exp(-4*(self.W@x-0.5))))*0.05,
                0, 1)
        return X.detach().numpy()


class OpenWormBaseline:
    def __init__(self):
        csv_text = None
        for p in ["/content/connectome.csv","connectome.csv"]:
            if os.path.exists(p):
                print(f"  OpenWorm: {p}...",end="",flush=True)
                csv_text=open(p).read().replace('\r\n','\n').replace('\r','\n')
                print(" OK"); break
        if csv_text is None:
            print("  OpenWorm: GitHub...",end="",flush=True)
            try:
                url=("https://raw.githubusercontent.com/adammarblestone-zz/"
                     "simple-C-elegans/master/OpenWorm/connectome.csv")
                csv_text=requests.get(url,timeout=30).text.replace('\r\n','\n').replace('\r','\n')
                print(" OK")
            except Exception as e:
                raise RuntimeError(f"connectome.csv bulunamadi: {e}")
        N=len(OW_NEURONS); idx={n:i for i,n in enumerate(OW_NEURONS)}
        conn=np.zeros((N,N)); gap=np.zeros((N,N)); gaba=[]
        for row in csv.reader(io.StringIO(csv_text)):
            if len(row)<4: continue
            pre,post,ct=row[0].strip(),row[1].strip(),row[2].strip()
            try: w=float(row[3])
            except: continue
            if pre not in idx or post not in idx: continue
            if ct=="Send":
                conn[idx[pre],idx[post]]+=w
                if len(row)>4 and "GABA" in row[4] and pre not in gaba:
                    gaba.append(pre)
            else: gap[idx[pre],idx[post]]+=w
        self.conn=conn; self.gap=gap; self.N=N; self.gaba=gaba
        self.ava=[idx[n] for n in ["AVAL","AVAR"] if n in idx]
        self.avb=[idx[n] for n in ["AVBL","AVBR"] if n in idx]
        Gcc=100e-12/10e-12; gsc=10e-12/10e-12; ggc=5e-12/10e-12
        Ec=np.full(N,-35e-3)
        rev=np.array([-45e-3 if OW_NEURONS[i] in gaba else 0.0 for i in range(N)])
        A=np.zeros((N,N)); b=np.zeros(N)
        for i in range(N):
            A[i,i]=1+(ggc/Gcc)*gap[:,i].sum()+(gsc/Gcc)*conn[:,i].sum()/2
            for j in range(N):
                if j!=i: A[i,j]=-(ggc/Gcc)*gap[j,i]
            b[i]=Ec[i]+(gsc/Gcc)*np.dot(conn[:,i],rev)/2
        self.Veq=np.linalg.solve(A,b)
        self.Ec=Ec; self.rev=rev; self.Gcc=Gcc; self.gsc=gsc
        self.ggc=ggc; self.Vr=35e-3; self.K=-4.39*8.0
        print(f"  OpenWorm: N={N}, syn={int((conn>0).sum())}, "
              f"gap={int((gap>0).sum())}")

    def _deriv(self,t,V):
        act=1/(1+np.exp(self.K*(V-self.Veq)/self.Vr))
        Ig=np.array([self.ggc*np.dot(self.gap[:,i],V[i]-V) for i in range(self.N)])
        Is=np.array([self.gsc*np.dot(self.conn[:,i],(V[i]-self.rev)*act)
                     for i in range(self.N)])
        return self.Gcc*(self.Ec-V)-Ig-Is

    def run(self,seed=42,T=2.0,n_steps=1200):
        rng=np.random.default_rng(seed)
        V0=self.Veq+rng.normal(0,2e-3,self.N)
        try:
            sol=solve_ivp(lambda t,V:self._deriv(t,V),[0,T],V0,
                          t_eval=np.linspace(0,T,n_steps),method="RK45",
                          max_step=T/n_steps*5,rtol=1e-4,atol=1e-6)
            return 1/(1+np.exp(-(sol.y.T-(-40e-3))/self.Vr))
        except: return np.zeros((n_steps,self.N))

# ══════════════════════════════════════════════════════════════════════
# KATO 2015 SPEKTRUMU
# ══════════════════════════════════════════════════════════════════════
def load_kato_spectrum():
    try:
        print("  Kato 2015...",end="",flush=True)
        r=requests.get("https://api.osf.io/v2/nodes/2395t/files/osfstorage/",
                       timeout=30)
        url=next(f["links"]["download"] for f in r.json()["data"]
                 if f["attributes"]["name"]=="WT_NoStim.mat")
        resp=requests.get(url,stream=True,timeout=300)
        chunks=[]; n=0
        for chunk in resp.iter_content(1024*1024):
            chunks.append(chunk); n+=len(chunk)
            if n%(20*1024*1024)==0:
                print(f"{n//1024//1024}M..",end="",flush=True)
        f=h5py.File(io.BytesIO(b"".join(chunks)),"r")
        traces=np.array(f["#refs#/Zk/traces"]); f.close()
        if traces.shape[0]<traces.shape[1]: traces=traces.T
        traces=scipy_detrend(traces,axis=0)
        traces=(traces-traces.mean(0))/(traces.std(0)+1e-8)
        nc=min(15,traces.shape[0],traces.shape[1])
        spec=PCA(n_components=nc).fit(traces).explained_variance_ratio_
        print(f" OK  PC1={spec[0]*100:.1f}%"); return spec
    except Exception as e:
        print(f"\n  Kato yuklenemedi ({e}) → sentetik")
        k=np.exp(-np.arange(15)*0.3); return k/k.sum()

# ══════════════════════════════════════════════════════════════════════
# YUKLEME
# ══════════════════════════════════════════════════════════════════════
def _find_connectome():
    for c in (["cook2019connectome.xlsx","/content/cook2019connectome.xlsx"]
               +glob.glob("/content/cook2019connectome.xlsx") + glob.glob("/content/*.xlsx")+glob.glob("*.xlsx")):
        if os.path.exists(c): return c
    return None

class OperatorPool(nn.Module):
    """
    Operator havuzu -- NeurosparkSNT B cekirdegi.

    INPUT (Spark)  : γ_f, γ_d
    SYSTEM         : γ_n, γ_c, γ_t, γ_learn
    BEHAVIOR       : γ_fwd, γ_rev, γ_turn
    TASK           : γ_DA, γ_lim, γ_cns, γ_5HT
    H              : Hebb coordination

    Phase dispatch (PHASE_OPS):
      ACTIVE_FWD  : {γ_f, γ_d, γ_n, γ_fwd, γ_t, H}
      ACTIVE_REV  : {γ_f, γ_d, γ_n, γ_rev, γ_t, H}
      ACTIVE_TURN : {γ_f, γ_d, γ_n, γ_turn, γ_t, H}  ← no FWD/REV conflict
      RESET       : {γ_c, γ_t}
      ACTIVE_TURN : {γ_f, γ_d, γ_n, γ_turn, γ_t, H}
      NEUTRAL     : {γ_f, γ_d, γ_n, γ_t, γ_c, H}

    γ_learn always runs (T_c side-effect only, no dx).
    γ_c excluded from ACTIVE phases by design (no hack needed).
    """

    SENS_THR = 0.25
    STIM_THR = 0.20

    # ── Dispatch tablosu ─────────────────────────────────────────────
    # Layer 1: Locomotion phase → base operators
    PHASE_OPS = {
        PHASE_ACTIVE_FWD:  {'γ_f','γ_d','γ_n','γ_fwd','γ_t','H'},
        PHASE_ACTIVE_REV:  {'γ_f','γ_d','γ_n','γ_rev','γ_t','H'},
        PHASE_RESET:       {'γ_c','γ_t'},
        PHASE_NEUTRAL:     {'γ_f','γ_d','γ_n','γ_t','γ_c','H'},
        PHASE_ACTIVE_TURN: {'γ_f','γ_d','γ_n','γ_turn','γ_t','H'},
    }

    # Katman 2: Chain task stage → ADD to PHASE_OPS
    # γ_DA, γ_lim, γ_cns, γ_5HT: task-specific operators
    TASK_ADD = {
        TASK_IDLE:     frozenset(),
        TASK_APPROACH: frozenset({'γ_DA','γ_lim'}),   # approach: dopamine + liminal signal
        TASK_CONTACT:  frozenset({'γ_cns','γ_5HT'}),  # temas: ye + serotonin
        TASK_RESET:    frozenset(),
    }

    # Katman 3: Chain task stage → REMOVE from PHASE_OPS (override)
    TASK_REMOVE = {
        TASK_IDLE:     frozenset(),
        TASK_APPROACH: frozenset(),
        TASK_CONTACT:  frozenset({'γ_fwd','γ_turn'}),
        TASK_RESET:    frozenset({'γ_fwd','γ_rev','γ_turn'}),
    }

    def compute_active_set(self, phase, task_stage):
        """Phase + gorev asamasindan active operator setini hesapla."""
        base = set(self.PHASE_OPS.get(phase, {'γ_f','γ_d','γ_n','γ_t','H'}))
        base |= self.TASK_ADD.get(task_stage, frozenset())
        base -= self.TASK_REMOVE.get(task_stage, frozenset())
        return base

    def __init__(self):
        super().__init__()
        # Same parameter set as A (comparability)
        self.g_fwd  = 1.50; self.g_rev  = 1.50; self.g_turn = 1.20
        self.gf     = 0.05; self.gn     = 1.50
        self.gd     = 0.50; self.gc     = 1.20
        self.gt     = 0.15; self.tau    = 0.80
        self.tauf   = 0.10; self.inh    = 0.45
        self.k_hebb = 0.15; self.zeta   = 0.15
        self.T_c         = 0.40   # cultivation temp memory (normalized 20°C)
        self.alpha_learn = 0.001  # T_c adaptation rate (slow, minute-scale)

    # ── INPUT (Spark) operatorleri ───────────────────────────────────

    def _gamma_f(self, N, phase, rng, dt):
        """γ_f Input(Spark): kesif gurultusu"""
        ns = 1.5 if phase==PHASE_RESET else 1.2 if phase==PHASE_NEUTRAL else 1.0
        return self.gf * ns * torch.randn(N, generator=rng) * np.sqrt(dt/self.tauf)

    def _gamma_d(self, x, sa_l, sr_l, food_l,
                 attract_raw, repel_eff, food_raw, dt):
        """γ_d Input(Spark): esikli uyaran amplifikasyonu"""
        d = torch.zeros(len(x))
        if attract_raw > self.STIM_THR and sa_l:
            d[sa_l] += self.gd * attract_raw * (1-x[sa_l]) * dt
        if repel_eff > self.STIM_THR and sr_l:
            d[sr_l] += self.gd * repel_eff * (1-x[sr_l]) * dt
        if food_raw > self.STIM_THR and sa_l:
            d[sa_l] += self.gd * 0.3 * food_raw * (1-x[sa_l]) * dt
        return d

    # ── SYSTEM operatorleri ───────────────────────────────────────────

    def _gamma_n(self, x, avb_l, ava_l,
                 attract_raw, repel_eff, tpost_raw, tant_raw,
                 temp_raw, o2_raw, food_raw, dt):
        """γ_n: karar bias -- AVB/AVA'ya esikli sensor sinyali"""
        d = torch.zeros(len(x))
        a_eff   = max(0.0, attract_raw - self.SENS_THR)
        r_eff   = max(0.0, repel_eff   - self.SENS_THR)
        tp_eff  = max(0.0, tpost_raw   - self.SENS_THR)
        ta_eff  = max(0.0, tant_raw    - self.SENS_THR)
        o2_eff  = max(0.0, o2_raw      - self.SENS_THR) * 0.4
        HIGH_TEMP_THR = 0.45
        temp_stress = max(0.0, temp_raw - HIGH_TEMP_THR) * 1.5
        temp_err    = temp_raw - self.T_c
        tc_fwd = max(0.0, -temp_err - 0.05) * 0.35
        tc_rev = max(0.0,  temp_err - 0.05) * 0.35
        fwd_bias = a_eff*0.4 + tp_eff*0.5 + tc_fwd
        rev_bias = r_eff*0.5 + ta_eff*0.6 + o2_eff + temp_stress + tc_rev
        food_bias_eff = max(0.0, food_raw - self.SENS_THR) * 0.30
        fwd_bias += food_bias_eff
        if avb_l and fwd_bias > 0:
            d[avb_l] = self.gn * fwd_bias * (1-x[avb_l]) * dt
        if ava_l and rev_bias > 0:
            d[ava_l] = self.gn * rev_bias * (1-x[ava_l]) * dt
        return d, fwd_bias, rev_bias

    def _gamma_c(self, x, theta, phase, turn_l, pause_l, dt):
        """γ_c: cyclic reset -- SADECE RESET ve NEUTRAL'da"""
        gc = 4.0 if phase==PHASE_RESET else self.gc * 0.3
        d = -gc * x * (x > theta).float() * dt
        if phase==PHASE_RESET and turn_l:
            d[turn_l] += self.gc * 0.4 * (1-x[turn_l]) * dt
        if phase==PHASE_NEUTRAL and pause_l:
            d[pause_l] += self.gc * 0.15 * (1-x[pause_l]) * dt
        return d

    def _gamma_t(self, x, steer_l, food_raw, temp_raw,
                 attract_raw, phase, dt):
        """γ_t: tonik homeostaz + dopamin + bas salinimi"""
        dopamin_brake = float(np.clip(food_raw * 0.5, 0, 0.35))
        d = self.gt * (x.mean()-x) * dt
        # Decay
        d += -self.gt * (3.0 if phase==PHASE_RESET else 1.0) * x * dt
        # Weathervane
        if steer_l:
            steer_drive = temp_raw if temp_raw > self.STIM_THR else attract_raw*0.3
            if steer_drive > 0.05:
                d[steer_l] += self.gt * steer_drive * (1-x[steer_l]) * dt * 0.6
        return d, dopamin_brake

    # ── BEHAVIOR operatorleri ─────────────────────────────────────────

    def _gamma_fwd(self, x, avb_l, fwd_m, ava_l, dopamin_brake, dt):
        """γ_fwd: FORWARD only -- hard lock"""
        d = torch.zeros(len(x))
        exc = self.g_fwd * (1-dopamin_brake)
        if avb_l: d[avb_l] = exc * (1-x[avb_l]) * dt
        if fwd_m:  d[fwd_m] = exc * 0.5 * (1-x[fwd_m]) * dt
        if ava_l:  d[ava_l] = -self.inh * x[ava_l] * dt
        return d

    def _gamma_rev(self, x, ava_l, rev_m, avb_l, dt):
        """γ_rev: REVERSE only -- hard lock"""
        d = torch.zeros(len(x))
        if ava_l: d[ava_l] = self.g_rev * (1-x[ava_l]) * dt
        if rev_m:  d[rev_m] = self.g_rev * 0.5 * (1-x[rev_m]) * dt
        if avb_l:  d[avb_l] = -self.inh * x[avb_l] * dt
        return d

    # ── TASK-SPECIFIC operatorleri ──────────────────────────────────

    def _gamma_DA(self, x, avb_l, sa_l, da_level, sim_attract, dt):
        """γ_DA: Dopamine -- approach motivation (APPROACH)
        Feeds AVB, stimulates attract neurons with sim_attract."""
        d = torch.zeros(len(x))
        if avb_l and da_level > 0:
            d[avb_l] += da_level * sim_attract * 0.4 * (1-x[avb_l]) * dt
        if sa_l and sim_attract > 0:
            d[sa_l]  += sim_attract * 0.3 * (1-x[sa_l]) * dt
        return d

    def _gamma_lim(self, x, avb_l, sim_attract, dt):
        """γ_lim: Liminal -- gradient tabanli yonelim
        Yaklastikca guclenen FWD bias."""
        d = torch.zeros(len(x))
        if avb_l and sim_attract > 0.1:
            d[avb_l] += sim_attract * 0.3 * (1-x[avb_l]) * dt
        return d

    def _gamma_cns(self, x, pharynx_l, dt):
        """γ_cns: Consume -- pharynx aktivasyonu (yeme motoru)"""
        d = torch.zeros(len(x))
        if pharynx_l:
            d[pharynx_l] += 1.5 * (1-x[pharynx_l]) * dt
        return d

    def _gamma_5HT(self, x, avb_l, pause_l, sht_level, dt):
        """γ_5HT: Serotonin -- doyum + lokomosyon yavaslatma
        AVB'yi bastirir, RIS'i aktive eder."""
        d = torch.zeros(len(x))
        if avb_l and sht_level > 0:
            d[avb_l] -= sht_level * 0.5 * x[avb_l] * dt    # yavasla
        if pause_l and sht_level > 0:
            d[pause_l] += sht_level * 0.4 * (1-x[pause_l]) * dt  # pause/satiety
        return d

    def _gamma_turn(self, x, turn_l, motor_turn_l, ava_l, avb_l, dt):
        """γ_turn: omega turn — PHASE_ACTIVE_TURN only.
        Activates circuit_turn (AIB/RIM) + motor_turn (SMD/RMD).
        No FWD/REV conflict: excluded from those PHASE_OPS sets.
        """
        d = torch.zeros(len(x))
        if turn_l:       d[turn_l]       += self.g_turn*(1-x[turn_l])*dt
        if motor_turn_l: d[motor_turn_l] += self.g_turn*0.6*(1-x[motor_turn_l])*dt
        if ava_l:        d[ava_l]        -= self.inh*x[ava_l]*dt
        if avb_l:        d[avb_l]        -= self.inh*0.5*x[avb_l]*dt
        return d

    def _gamma_learn(self, learn_l, temp_raw, x, dt):
        """γ_learn: T_c dynamic via circuit_learn (AIY/AIZ).
        Side-effect only — no dx contribution.
        """
        if not learn_l or float(temp_raw) <= 0.10:
            return
        la = float(x[learn_l].mean())
        self.T_c = float(np.clip(
            self.T_c + self.alpha_learn*(float(temp_raw)-self.T_c)*la*dt,
            0.10, 0.90))

    def _hebb(self, x, H, dt):
        """H: Hebb coordination"""
        return self.k_hebb * torch.tanh(H@x) * dt

    def _connectome(self, x, W, dt):
        """Connectome dinamigi"""
        return (1/self.tau)*(-x + torch.tanh(W@x))*dt

    def forward(self, x, theta, W, H,
                ava_l, avb_l, fwd_m, rev_m,
                steer_l, turn_l, pause_l,
                sa_l, sr_l, tant_l, tpost_l,
                food_l, temp_l, co2_l, o2_l,
                pharynx_l,
                phase, task_stage, rng, dt=0.05,
                da_level=0.0, sht_level=0.0, sim_attract=0.0,
                learn_l=None, motor_turn_l=None,
                has_consume=True):
        # task_stage: from ChainTask. has_consume: skip γ_cns/γ_5HT if yemek_bul
        N = len(x)

        # ── Sensor on-hesaplama (ortak, tek seferlik) ─────────────────
        attract_raw = float(x[sa_l].mean())    if sa_l    else 0.0
        repel_raw   = float(x[sr_l].mean())    if sr_l    else 0.0
        co2_raw     = float(x[co2_l].mean())   if co2_l   else 0.0
        o2_raw      = float(x[o2_l].mean())    if o2_l    else 0.0
        tant_raw    = float(x[tant_l].mean())  if tant_l  else 0.0
        tpost_raw   = float(x[tpost_l].mean()) if tpost_l else 0.0
        food_raw    = float(x[food_l].mean())  if food_l  else 0.0
        temp_raw    = float(x[temp_l].mean())  if temp_l  else 0.0
        repel_eff   = repel_raw + co2_raw * 0.5

        # ── Active operator set: phase + task dispatch ────────────────
        active = self.compute_active_set(phase, task_stage)

        dx = torch.zeros(N)

        # Connectome her zaman (temel dinamik)
        dx += self._connectome(x, W, dt)

        if 'γ_f' in active:
            dx += self._gamma_f(N, phase, rng, dt)

        if 'γ_d' in active:
            dx += self._gamma_d(x, sa_l, sr_l, food_l,
                                attract_raw, repel_eff, food_raw, dt)

        fwd_bias = 0.0; rev_bias = 0.0; dopamin_brake = 0.0
        if 'γ_n' in active:
            d_n, fwd_bias, rev_bias = self._gamma_n(
                x, avb_l, ava_l, attract_raw, repel_eff,
                tpost_raw, tant_raw, temp_raw, o2_raw, food_raw, dt)
            dx += d_n

        if 'γ_c' in active:
            dx += self._gamma_c(x, theta, phase, turn_l, pause_l, dt)

        if 'γ_t' in active:
            d_t, dopamin_brake = self._gamma_t(
                x, steer_l, food_raw, temp_raw, attract_raw, phase, dt)
            dx += d_t

        if 'γ_fwd' in active:
            dx += self._gamma_fwd(x, avb_l, fwd_m, ava_l, dopamin_brake, dt)

        if 'γ_rev' in active:
            dx += self._gamma_rev(x, ava_l, rev_m, avb_l, dt)

        if 'H' in active:
            dx += self._hebb(x, H, dt)

        # ── Task-specific operatorler (chain state'ten gelir) ─────────
        if 'γ_DA' in active:
            dx += self._gamma_DA(x, avb_l, sa_l, da_level, sim_attract, dt)

        if 'γ_lim' in active:
            dx += self._gamma_lim(x, avb_l, sim_attract, dt)

        # γ_cns + γ_5HT: only for yemek_ye (has_consume=True)
        # yemek_bul: CONTACT tagged only, no pharynx, no 5HT slowing
        if 'γ_cns' in active and has_consume:
            dx += self._gamma_cns(x, pharynx_l, dt)

        if 'γ_5HT' in active and has_consume:
            dx += self._gamma_5HT(x, avb_l, pause_l, sht_level, dt)

        if 'γ_turn' in active:
            dx += self._gamma_turn(x, turn_l, motor_turn_l, ava_l, avb_l, dt)

        # γ_learn: always active (side-effect only)
        self._gamma_learn(learn_l, temp_raw, x, dt)

        return dx, {
            "attract":       attract_raw,
            "repel":         repel_eff,
            "food":          food_raw,
            "dopamin_brake": dopamin_brake,
            "fwd_bias":      fwd_bias,
            "rev_bias":      rev_bias,
            "active_ops":    sorted(active),
            "task_stage":    task_stage,
        }


class NeuroSparkSNT_B(nn.Module):
    """
    NeuroSparkSNT B -- Operator Pool Architecture

    NeurosparkSNT A ile ayni parameterler, farkli mimari.
    Phase-based operator dispatch: unnecessary operators do not run.
    No gc_eff=0 hack -- γ_c is simply absent from ACTIVE_FWD pool.
    """
    def __init__(self, W_circuit, registry, dt=0.05):
        super().__init__()
        self.reg   = registry
        self.dim_n = W_circuit.shape[0]
        self.dt    = dt
        self.bnames = list(BEHAVIOR_CIRCUITS.keys())

        W_t = torch.tensor(W_circuit, dtype=torch.float32)
        self.register_buffer("W", W_t)
        self.register_buffer("H", (W_t+W_t.T)/2.0)

        li = registry.local_idx
        self.ava_l   = li("circuit_reverse")[:2]
        self.avb_l   = li("circuit_forward")[:2]
        self.fwd_m   = li("motor_forward")
        self.rev_m   = li("motor_reverse")
        self.steer_l = li("circuit_steer")
        self.turn_l  = li("circuit_turn")
        self.pause_l = li("circuit_pause")
        self.sa_l    = li("sensory_attract")
        self.sr_l    = li("sensory_repel")
        self.tant_l  = li("sensory_touch_ant")
        self.tpost_l = li("sensory_touch_post")
        self.food_l  = li("sensory_food")
        self.temp_l  = li("sensory_temp")
        self.co2_l   = li("sensory_co2")
        self.o2_l    = li("sensory_oxygen")
        self.pharynx_l   = li("pharynx")
        self.learn_l     = li("circuit_learn")
        self.motor_turn_l= li("motor_turn")

        self.stim        = _StimulusLayer(self.dim_n, registry)
        self.pool        = OperatorPool()          # A: ops, B: pool
        self.cmd         = _CommandLayer(registry)
        self.clamp       = _SoftClamp()
        self.phase_ctrl  = PhaseController(dt=dt)
        self.dispatcher  = TaskDispatcher()
        self.chain       = ChainTaskPipeline(registry, self.dispatcher)
        self.awareness   = StimulusAwarenessLayer()

        print(f"  NeuroSparkSNT_B (Pool): dim={self.dim_n}")
        print(f"  AVA={self.ava_l}  AVB={self.avb_l}")
        print(f"  OperatorPool: phase-based dispatch active")

    def _fresh(self, rng):
        self.phase_ctrl._rng = np.random.default_rng(
            int(torch.randint(0,99999,(1,),generator=rng).item()))
        self.phase_ctrl.reset()
        self.chain.reset()
        return {
            "x":     torch.rand(self.dim_n, generator=rng)*0.15+0.10,
            "theta": torch.full((self.dim_n,), self.clamp.init),
            "vt":    torch.zeros(self.dim_n),
            "vx":    torch.zeros(self.dim_n),
        }

    def step(self, state, stype="none", t=0.0, rng=None):
        if rng is None: rng = torch.Generator()
        x, theta, vt, vx = (state["x"], state["theta"],
                             state["vt"], state["vx"])

        ava_a = float(x[self.ava_l].mean())  if self.ava_l  else 0.0
        avb_a = float(x[self.avb_l].mean())  if self.avb_l  else 0.0
        rep_a = float(x[self.sr_l].mean())   if self.sr_l   else 0.0
        tan_a = float(x[self.tant_l].mean()) if self.tant_l else 0.0

        if self.dispatcher.is_goal_directed(stype) and self.chain.stage == TASK_IDLE:
            self.chain.activate(stype)
        effective_stype = (stype if not self.dispatcher.is_goal_directed(stype)
                           else "food_odor")

        self.awareness.update(effective_stype)
        phase, _ = self.phase_ctrl.update(ava_a, avb_a, rep_a, tan_a,
                                           awareness=self.awareness)

        stim = self.stim(x, effective_stype, t=t)

        # ── Model B: Chain state ONCE guncellenir, pool SONRA dispatch yapar ──
        food_raw_direct = float(x[self.food_l].mean()) if self.food_l else 0.0
        _contact, _stage = self.chain.advance_state(phase, food_raw_direct, self.dt)
        ps = self.chain.get_pool_state()  # da_level, sht_level, sim_attract

        dx, info = self.pool(
            x, theta, self.W, self.H,
            self.ava_l, self.avb_l, self.fwd_m, self.rev_m,
            self.steer_l, self.turn_l, self.pause_l,
            self.sa_l, self.sr_l, self.tant_l, self.tpost_l,
            self.food_l, self.temp_l, self.co2_l, self.o2_l,
            self.pharynx_l,                            # task-specific
            phase, ps["stage"], rng, self.dt,          # dispatch: phase + task
            da_level     = ps["da_level"],
            sht_level    = ps["sht_level"],
            sim_attract  = ps["sim_attract"],
            learn_l      = self.learn_l,
            motor_turn_l = self.motor_turn_l,
            has_consume  = ps.get('has_consume', True))
        dx = dx + stim*self.dt

        chain_info = {**ps, "progress": self.chain.lim.progress}

        vx_new = (1-self.pool.zeta)*dx + self.pool.zeta*vx
        x_new  = torch.clamp(x+vx_new, 0.0, 1.0)

        with torch.no_grad():
            x_new = x_new + torch.randn_like(x_new) * 1e-5
            x_new = torch.clamp(x_new, 0.0, 1.0)

        beh    = self.cmd(x_new, phase)
        th_new, vt_new = self.clamp(
            x_new, theta, vt, phase, info["attract"], info["repel"])

        return {"x":x_new, "theta":th_new, "vt":vt_new, "vx":vx_new,
                "beh":beh, "phase":phase, "info":info, "chain":chain_info}

    def simulate(self, n_steps=1200, stype="food_odor",
                 stim_on=5.0, stim_off=15.0, seed=42, verbose=True):
        rng = torch.Generator(); rng.manual_seed(seed)
        torch.manual_seed(seed)
        t_arr = np.arange(n_steps)*self.dt
        state = self._fresh(rng)
        X   = torch.zeros(n_steps, self.dim_n)
        B   = torch.zeros(n_steps, len(self.bnames))
        ph  = np.zeros(n_steps, dtype=int)
        chain_stage_hist = np.zeros(n_steps, dtype=int)

        if verbose:
            print(f"\n{'='*68}")
            print(f"NeuroSparkSNT_B  stype={stype}  [{stim_on}-{stim_off}s]")
            print(f"{'='*68}")

        for i, t in enumerate(t_arr):
            st = stype if stim_on<=t<=stim_off else "none"
            out = self.step(state, st, t, rng)
            state = out
            X[i]=out["x"]; B[i]=out["beh"]; ph[i]=out["phase"]
            chain_stage_hist[i] = out["chain"]["stage"]

        ava_avb_r = float('nan')
        if self.ava_l and self.avb_l:
            a = X[:,self.ava_l].mean(1).detach().numpy()
            b = X[:,self.avb_l].mean(1).detach().numpy()
            if a.std()>1e-6 and b.std()>1e-6:
                ava_avb_r, _ = pearsonr(a, b)

        pp = {PHASE_NAMES[p]:float((ph==p).mean()*100) for p in range(5)}
        dom_b = self.bnames[B.mean(0).argmax().item()]

        if verbose:
            print(f"  AVA-AVB r: {ava_avb_r:+.4f}  Dominant: {dom_b}")
            print(f"  Phases: "+" ".join(f"{k}:{v:.0f}%" for k,v in pp.items()))

        return {"X":X,"B":B,"ava_avb_r":ava_avb_r,
                "dom":dom_b,"phase_pcts":pp,
                "chain_stage_hist":chain_stage_hist, "ph": ph}
print("NeuroSparkSNT -- Symplectic Neural Topology")
print("="*52)

_cp = _find_connectome()
if _cp is None:
    print("HATA: Connectome.xlsx bulunamadi!")
    print("  from google.colab import files; files.upload()")
else:
    print(f"Connectome: {_cp}")

    print("\n[1/5] Cook 2019...")
    W_cook, names = load_cook2019(_cp)
    _rt = NeuronRegistry(names)
    W_cook = add_reciprocal_inhibition(W_cook, _rt.ava, _rt.avb, strength=0.8)

    print("\n[2/5] NeuroSparkSNT...")
    reg = NeuronRegistry(names)
    reg.report()
    W_circ = reg.get_circuit_W(W_cook)
    print(f"  Circuit: {W_circ.shape[0]} neuron")
    ns_snt = NeuroSparkSNT(W_circ, reg)

    # Alias for benchmark
    reg_full = reg

    print("\n[3/5] Wilson-Cowan...")
    wc = WilsonCowanBaseline(W_cook)

    print("\n[4/5] OpenWorm HH-ODE...")
    try: ow=OpenWormBaseline(); ow_ok=True
    except Exception as e:
        print(f"  UYARI: {e}"); ow=None; ow_ok=False

    print("\n[5/5] Kato 2015 PCA spektrumu...")
    kato_spec = load_kato_spectrum()

    print("\n[+] NeuroSparkSNT_B (Pool architecture)...")
    ns_snt_b = NeuroSparkSNT_B(W_circ, reg)

    print("\n✓ Ready: ns_snt, ns_snt_b, wc, ow, kato_spec, reg, reg_full")
    print("  exec(open('neurosparksnt_bench.py').read())")

# ══════════════════════════════════════════════════════════════════════
# NeuroSparkSNT B -- Operator Pool Architecture
# ══════════════════════════════════════════════════════════════════════
# Fark: Her operator ayri metod. Phase/task asamasina gore havuzdan cekilir.
# In A: all operators run every step (gc_eff=0 hack to disable some)
# In B: only required operators are selected via set-dispatch
#
# Benchmark katkisi:
#   Same biological parameter set → measures effect of architectural difference
#   A kazanirsa: monolith sinerji onemli
#   B kazanirsa: pasimon + modulerlik ustun

