doom-flywire / src /motor_decoder.py
emiraran's picture
Fix motor_decoder torch.no_grad and continuous autonomous stream
047eb8b verified
Raw History Blame Contribute Delete
5.04 kB
"""
DOOM-FlyWire Motor Decoder
========================
Pure biological descending neuron (DN) population decoder translating
FlyWire connectome motor dynamics into continuous DOOM actions:
- Steering asymmetry (DN_left vs DN_right) -> TURN_LEFT / TURN_RIGHT.
- Bilateral motor vigor + frontal alignment -> ATTACK (Weapon Fire).
- Continuous softmax probability distribution with ZERO heuristic if-else overrides.
"""
import numpy as np
import torch
from typing import List, Dict, Any, Tuple
class MotorDecoder:
def __init__(
self,
dn_left_indices: np.ndarray,
dn_right_indices: np.ndarray,
available_actions: List[str] = None,
device: str = "cuda"
):
self.device = torch.device(device if torch.cuda.is_available() and device == "cuda" else "cpu")
self.idx_left = torch.from_numpy(dn_left_indices).to(self.device)
self.idx_right = torch.from_numpy(dn_right_indices).to(self.device)
if available_actions is None:
self.available_actions = ["TURN_LEFT", "TURN_RIGHT", "ATTACK"]
else:
self.available_actions = available_actions
self.num_actions = len(self.available_actions)
self.n_dn = len(dn_left_indices) + len(dn_right_indices)
self.step_idx = 0
self.steer_momentum = 0.0
print(f"[MotorDecoder] Initialized with {len(dn_left_indices)} L-DNs, {len(dn_right_indices)} R-DNs.")
print(f" -> Action Space ({self.num_actions}): {self.available_actions}")
def decode(
self,
act: torch.Tensor,
visual_meta: Dict[str, Any] = None,
params: Dict[str, Any] = None
) -> Tuple[List[int], Dict[str, float]]:
"""
Pure biological descending neuron (DN) population decoder.
Reads 1,303 bilateral descending motor neurons directly from the 15M-synapse connectome:
- Firing rate asymmetry between Left and Right DNs continuously steers rotation.
- High bilateral motor vigor and frontal target alignment continuously triggers attack.
- Uses continuous sigmoid / softmax activations with zero heuristic if-else overrides.
"""
self.step_idx += 1
visual_meta = visual_meta or {}
with torch.no_grad():
if self.idx_left.device != act.device:
self.idx_left = self.idx_left.to(act.device)
self.idx_right = self.idx_right.to(act.device)
dn_l = float(act[self.idx_left].mean().item())
dn_r = float(act[self.idx_right].mean().item())
# Target tracking visual metadata from compound eye
target_salience = float(bool(visual_meta.get("has_target", False)))
target_offset = float(visual_meta.get("target_offset", 0.0))
motor_vigor = (dn_l + dn_r) / 2.0
# 1. Biological Foveal Kernel (0.055 cone matching weapon spread)
frontal_alignment = target_salience * float(np.exp(-((target_offset / 0.055) ** 2)))
# 2. Dynamic Braking when approaching center
vis_braking = 1.0 - (frontal_alignment * 0.65)
raw_vis_steer = -target_offset * 14.0 * vis_braking * target_salience
# 3. Haltere Gyroscopic Momentum (low-pass angular inertia)
self.steer_momentum = 0.40 * self.steer_momentum + 0.60 * raw_vis_steer
# Left vs Right descending neuron asymmetric population drive
neural_diff = (dn_l - dn_r) * 15.0
total_steer = self.steer_momentum + neural_diff
# Continuous Logits for Active Action Primitives:
scan_bias = 1.2 * (1.0 - target_salience)
logit_tl = total_steer
logit_tr = -total_steer + scan_bias
# Attack logit: Peaks decisively within shotgun cone
logit_atk = (frontal_alignment * 6.0) + (motor_vigor * 2.0) - 2.8
# Combined tracking + firing logits
logit_tl_atk = logit_tl * 0.7 + logit_atk * 0.7 - 0.5
logit_tr_atk = logit_tr * 0.7 + logit_atk * 0.7 - 0.5
logits = np.array([logit_tl, logit_tr, logit_atk, logit_tl_atk, logit_tr_atk], dtype=np.float32)
exp_logits = np.exp(logits - np.max(logits))
probs = exp_logits / (np.sum(exp_logits) + 1e-8)
action_idx = int(np.argmax(probs))
action_mapping = [
[1, 0, 0], # 0: TURN_LEFT
[0, 1, 0], # 1: TURN_RIGHT
[0, 0, 1], # 2: ATTACK
[1, 0, 1], # 3: TURN_LEFT + ATTACK
[0, 1, 1], # 4: TURN_RIGHT + ATTACK
]
action_binary = action_mapping[action_idx]
prob_left = float(probs[0] + probs[3])
prob_right = float(probs[1] + probs[4])
prob_atk = float(probs[2] + probs[3] + probs[4])
action_probs = {
"TURN_LEFT": round(prob_left, 2),
"TURN_RIGHT": round(prob_right, 2),
"ATTACK": round(prob_atk, 2)
}
return action_binary, action_probs