""" Step 6 acceptance tests: forward-transport visualization, firewalled. (A) interference visible -> inject 1 mode; energy spreads across modes/layers (B) energy conserved -> forward trace conserves total intensity per layer (C) FIREWALL -> rendering has zero effect on the compute path: logits before == after, and the grid carries no grad """ import torch from photonic.mzi import MZIMesh from photonic.hybrid import HybridPhotonicNet from photonic.visualize import render_energy_svg torch.manual_seed(0) def test_interference_visible(n=8): mesh = MZIMesh(n, seed=7, phase_bits=6) x = torch.zeros(n, dtype=torch.complex64) x[0] = 1.0 # inject ONE mode grid = mesh.forward_trace(x) render_energy_svg(grid, "core_energy.svg") lit_out = int((grid[-1] > 1e-3).sum()) print(f"(A) interference: 1 input mode -> {lit_out}/{n} output modes lit; " f"wrote core_energy.svg -> {'PASS' if lit_out > 1 else 'FAIL'}") return lit_out > 1 def test_energy_conserved(n=8): mesh = MZIMesh(n, seed=7, phase_bits=6) x = torch.randn(n, dtype=torch.complex64) grid = mesh.forward_trace(x) per_layer = grid.sum(dim=1) # total intensity at each layer spread = float(per_layer.max() - per_layer.min()) print(f"(B) energy conserved through forward trace: layer totals vary by " f"{spread:.2e} -> {'PASS' if spread < 1e-4 else 'FAIL'}") return spread < 1e-4 def test_firewall(): net = HybridPhotonicNet(d_in=2, modes=8, n_classes=2, phase_bits=6) x = torch.randn(4, 2) before = net(x).detach().clone() # visualize using a detached forward trace of the optical core field = net.enc(x[:1]).to(torch.complex64) grid = net.optical.V.forward_trace(field[0]) assert not grid.requires_grad render_energy_svg(grid, "hybrid_core_energy.svg") after = net(x).detach().clone() unchanged = torch.equal(before, after) no_grad = not grid.requires_grad print(f"(C) FIREWALL: logits identical after render = {unchanged}; " f"grid carries no grad = {no_grad} -> {'PASS' if unchanged and no_grad else 'FAIL'}") return unchanged and no_grad if __name__ == "__main__": print("=" * 60) print("STEP 6 -- forward-transport visualization (firewalled)") print("=" * 60) results = [test_interference_visible(), test_energy_conserved(), test_firewall()] print("-" * 60) print(f"RESULT: {sum(results)}/{len(results)} acceptance tests passed")