File size: 26,583 Bytes
3579fb4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0654a8d
3579fb4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0654a8d
3579fb4
 
0654a8d
 
 
 
3579fb4
 
0654a8d
 
 
 
3579fb4
 
0654a8d
 
 
 
 
 
3579fb4
 
0654a8d
 
 
3579fb4
 
0654a8d
 
3579fb4
 
0654a8d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3579fb4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0654a8d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3579fb4
 
 
 
0654a8d
3579fb4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0654a8d
3579fb4
 
 
 
 
 
 
 
 
 
 
 
 
 
0654a8d
 
3579fb4
 
0654a8d
3579fb4
 
 
 
 
0654a8d
 
3579fb4
0654a8d
 
 
 
 
3579fb4
0654a8d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3579fb4
0654a8d
 
3579fb4
0654a8d
 
 
 
 
 
 
 
 
 
 
 
 
3579fb4
0654a8d
 
 
 
3579fb4
 
0654a8d
 
3579fb4
0654a8d
 
 
 
3579fb4
 
0654a8d
 
 
 
 
3579fb4
0654a8d
 
 
3579fb4
0654a8d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3579fb4
 
0654a8d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3579fb4
0654a8d
 
3579fb4
0654a8d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3579fb4
0654a8d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
"""Exact self-reproduction for the SUBLEQ threshold host.

This module supplies the parts the universal constructor needs in order to
emit its own complete instance and not only its weight file: a tape device
with a rewind request and an end-of-tape status, the three-phase program P*,
the framing maps ser / inst, and runners for the three evaluators.

Device (memory-mapped, all logic in the runtime, none of it in the netlist):

    0xF9  C_WR    write request     writing 1 emits R_OUT, then R_OUT <- 0
    0xFA  C_EOT   end-of-tape       status maintained by the device
    0xFB  C_RW    rewind request    writing 1 sets the head to 0
    0xFC  C_RD    read request      writing 1 loads R_IN and advances the head
    0xFD  R_IN    input register
    0xFE  R_OUT   output register
    0xFF          halt program counter

One step is: execute the SUBLEQ instruction, then apply the device to the
resulting state in the order (write, read, rewind, clear requests). A request
fires on the value 1, not on the fact that the cell was addressed, so the
device reads the machine state and nothing else.

Program variables:

    0xF0  Z       constant 0 (restored by every instruction that uses it)
    0xF1  ONE     constant 1
    0xF2  T1      scratch: repeat counter
    0xF3  T2      scratch: literal counter
    0xF4  NEG1    constant 0xFF
    0xF5  EOK     constant 1, the end-of-tape test target
"""

from __future__ import annotations

import hashlib
import os
import struct
import sys
from typing import Dict, List, Optional, Tuple

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
HOST_PATH = os.path.join(REPO, "variants", "neural_subleq8io_netlist.safetensors")

# device cells
C_WR, C_EOT, C_RW, C_RD, R_IN, R_OUT = 0xF9, 0xFA, 0xFB, 0xFC, 0xFD, 0xFE
HALT_PC = 0xFF

# program variables
Z, ONE, T1, T2, NEG1, EOK = 0xF0, 0xF1, 0xF2, 0xF3, 0xF4, 0xF5


# =============================================================================
# Recipe language (grammar and decoder), reused from constructor8
# =============================================================================

def describe(data: bytes) -> bytes:
    """Compile bytes into a recipe. Literal tokens carry up to 127 bytes; a run
    of at least 4 equal bytes becomes a repeat token. Every byte is stored
    negated mod 256 so the machine recovers it with one subtraction."""
    tape = bytearray()
    i, n = 0, len(data)
    while i < n:
        j = i
        while j < n and data[j] == data[i] and j - i < 127:
            j += 1
        if j - i >= 4:
            tape.append(256 - (j - i))
            tape.append((256 - data[i]) % 256)
            i = j
            continue
        k = i
        while k < n and k - i < 127:
            m = k
            while m < n and data[m] == data[k] and m - k < 4:
                m += 1
            if m - k >= 4:
                break
            k += 1
        k = max(k, i + 1)
        tape.append(k - i)
        tape.extend((256 - x) % 256 for x in data[i:k])
        i = k
    tape.append(0)
    return bytes(tape)


def describe_literal(data: bytes) -> bytes:
    """The all-literal encoding used in the proof of the length bound."""
    tape = bytearray()
    for i in range(0, len(data), 127):
        block = data[i:i + 127]
        tape.append(len(block))
        tape.extend((256 - x) % 256 for x in block)
    tape.append(0)
    return bytes(tape)


def decode(tape: bytes) -> bytes:
    """delta: the decoding map on well-formed recipes."""
    out = bytearray()
    i = 0
    while True:
        t = tape[i]
        i += 1
        if t == 0:
            return bytes(out)
        if t <= 127:
            for _ in range(t):
                out.append((256 - tape[i]) % 256)
                i += 1
        elif t == 128:
            raise ValueError("tag 128 is reserved")
        else:
            out.extend([(256 - tape[i]) % 256] * (256 - t))
            i += 1


# =============================================================================
# Framing: ser and inst
# =============================================================================

def lam(k: int) -> bytes:
    """Eight-byte little-endian length field."""
    return struct.pack("<Q", k)


def field(u: bytes) -> bytes:
    return lam(len(u)) + u


def ser(sigma: bytes, m: bytes, tau: bytes) -> bytes:
    assert len(m) == 256
    return field(sigma) + field(m) + tau


def inst(s: bytes) -> Tuple[bytes, bytes, bytes]:
    """Partial inverse of ser: recover (sigma, m, tau)."""
    if len(s) < 8:
        raise ValueError("truncated")
    a = struct.unpack("<Q", s[:8])[0]
    if len(s) < 8 + a + 8:
        raise ValueError("truncated")
    sigma = s[8:8 + a]
    b = struct.unpack("<Q", s[8 + a:16 + a])[0]
    if b != 256:
        raise ValueError("memory image is not 256 bytes")
    m = s[16 + a:16 + a + b]
    if len(m) != b:                      # declared 256 bytes, fewer present
        raise ValueError("truncated memory image")
    tau = s[16 + a + b:]
    return sigma, m, tau


# =============================================================================
# Programs
# =============================================================================

def _emit(prog: List[Tuple[int, int, int]]) -> Dict[int, int]:
    mem = {}
    for idx, (a, b, c) in enumerate(prog):
        mem[idx * 3] = a
        mem[idx * 3 + 1] = b
        mem[idx * 3 + 2] = c
    return mem


# The decoding loop. Addresses are 3k; every branch target below is written as
# an instruction index and resolved to 3*index by _emit.
#
#  k0  T2 <- 0
#  k1  T1 <- 0
#  k2  request read of the tag
#  k3  T1 <- -T
#  k4  T2 <- T
#  k5  branch to k7 when T = 0 or T >= 128
#  k6  goto LITERAL
#  k7  branch to END when 256-T <= 0, that is T in {0,128}
#  k8  goto REPEAT with T1 = 256-T the run length
#  k9  LITERAL: request read of the next byte
#  k10 R_OUT <- b
#  k11 emit
#  k12 T2 <- T2-1; branch to k14 when the count is exhausted
#  k13 goto k9
#  k14 goto k0
#  k15 REPEAT: request read of the value byte
#  k16 R_OUT <- b
#  k17 emit
#  k18 T1 <- T1-1; branch to k14 when the count is exhausted
#  k19 goto k16
#  k20 END
DECODE_LOOP = [
    (T2, T2, 3 * 1),      # k0
    (T1, T1, 3 * 2),      # k1
    (NEG1, C_RD, 3 * 3),  # k2
    (R_IN, T1, 3 * 4),    # k3
    (T1, T2, 3 * 5),      # k4
    (Z, T2, 3 * 7),       # k5
    (Z, Z, 3 * 9),        # k6
    (Z, T1, 3 * 20),      # k7
    (Z, Z, 3 * 15),       # k8
    (NEG1, C_RD, 3 * 10),  # k9
    (R_IN, R_OUT, 3 * 11),  # k10
    (NEG1, C_WR, 3 * 12),  # k11
    (ONE, T2, 3 * 14),    # k12
    (Z, Z, 3 * 9),        # k13
    (Z, Z, 3 * 0),        # k14
    (NEG1, C_RD, 3 * 16),  # k15
    (R_IN, R_OUT, 3 * 17),  # k16
    (NEG1, C_WR, 3 * 18),  # k17
    (ONE, T1, 3 * 14),    # k18
    (Z, Z, 3 * 16),       # k19
]

# P: the constructor. The end token halts.
P = DECODE_LOOP + [(Z, Z, HALT_PC)]

# P*: the end token enters the rewind phase, then the copy phase.
#  k20 R:  request rewind
#  k21 B0: request read
#  k22 B1: EOK <- 1-EOT; halt when the end of tape is reached
#  k23 B2: Z <- -b
#  k24 B3: R_OUT <- b
#  k25 B4: emit
#  k26 B5: Z <- 0 and go to B0
P_STAR = DECODE_LOOP + [
    (NEG1, C_RW, 3 * 21),   # k20
    (NEG1, C_RD, 3 * 22),   # k21
    (C_EOT, EOK, HALT_PC),  # k22
    (R_IN, Z, 3 * 24),      # k23
    (Z, R_OUT, 3 * 25),     # k24
    (NEG1, C_WR, 3 * 26),   # k25
    (Z, Z, 3 * 21),         # k26
]


# P_e: P with an end-of-tape guard after the tag read, so that the machine
# halts on every tape, recipe or not. Instruction j3 assigns
# M[EOK] <- M[EOK] - M[C_EOT]; a successful tag read has cleared C_EOT and
# leaves EOK at 1, while a read at the end of the tape sets it and the result
# 0 transfers control to the halt cell.
P_TOTAL = [
    (T2, T2, 3 * 1),        # j0
    (T1, T1, 3 * 2),        # j1
    (NEG1, C_RD, 3 * 3),    # j2  read the tag
    (C_EOT, EOK, HALT_PC),  # j3  halt if that read was past the end of the tape
    (R_IN, T1, 3 * 5),      # j4
    (T1, T2, 3 * 6),        # j5
    (Z, T2, 3 * 8),         # j6
    (Z, Z, 3 * 10),         # j7
    (Z, T1, 3 * 21),        # j8
    (Z, Z, 3 * 16),         # j9
    (NEG1, C_RD, 3 * 11),   # j10 LITERAL
    (R_IN, R_OUT, 3 * 12),  # j11
    (NEG1, C_WR, 3 * 13),   # j12
    (ONE, T2, 3 * 15),      # j13
    (Z, Z, 3 * 10),         # j14
    (Z, Z, 3 * 0),          # j15
    (NEG1, C_RD, 3 * 17),   # j16 REPEAT
    (R_IN, R_OUT, 3 * 18),  # j17
    (NEG1, C_WR, 3 * 19),   # j18
    (ONE, T1, 3 * 15),      # j19
    (Z, Z, 3 * 17),         # j20
    (Z, Z, HALT_PC),        # j21 END
]


def decode_any(tape: bytes, guard: bool) -> Tuple[bytes, bool]:
    """The output of the decoding loop on an arbitrary tape.

    `guard` selects P_e over P. A read at the end of the tape sets the
    end-of-tape status and leaves the input register holding the last byte it
    received, so a token whose payload overruns the tape is completed with
    copies of that byte. Without the guard the loop diverges exactly when the
    tape is exhausted at a tag read and the stale input register holds neither
    0 nor 128; the second component of the result records whether the machine
    halts.
    """
    out = bytearray()
    h, rin = 0, 0
    while True:
        if h < len(tape):
            rin = tape[h]
            h += 1
            eot = 0
        else:
            eot = 1
        if guard and eot:
            return bytes(out), True
        t = rin
        if not guard and eot:
            return bytes(out), t in (0, 128)
        if t == 0 or t == 128:
            return bytes(out), True
        reps = 1 if t > 128 else t
        count = (256 - t) if t > 128 else 1
        for _ in range(reps):
            if h < len(tape):
                rin = tape[h]
                h += 1
            out.extend([(256 - rin) % 256] * count)


def memory_image(prog: List[Tuple[int, int, int]]) -> List[int]:
    """The 256-byte initial memory image holding a program and its constants."""
    mem = [0] * 256
    for addr, val in _emit(prog).items():
        mem[addr] = val & 0xFF
    mem[Z] = 0
    mem[ONE] = 1
    mem[T1] = 0
    mem[T2] = 0
    mem[NEG1] = 0xFF
    mem[EOK] = 1
    return mem


M_P = memory_image(P)
M_STAR = memory_image(P_STAR)


# =============================================================================
# Device
# =============================================================================

class Tape:
    """Environment state (tau, h, omega) with the operations of the definition."""

    def __init__(self, tau: bytes):
        self.tau = tau
        self.h = 0
        self.out = bytearray()

    def apply(self, mem: List[int]) -> None:
        """One device application to the post-instruction memory image."""
        if mem[C_WR] == 1:
            self.out.append(mem[R_OUT])
            mem[R_OUT] = 0
        if mem[C_RD] == 1:
            if self.h < len(self.tau):
                mem[R_IN] = self.tau[self.h]
                mem[C_EOT] = 0
                self.h += 1
            else:
                mem[C_EOT] = 1
        if mem[C_RW] == 1:
            self.h = 0
        mem[C_WR] = 0
        mem[C_RD] = 0
        mem[C_RW] = 0


# =============================================================================
# Evaluator 1: integer reference
# =============================================================================

def run_reference(mem0: List[int], tau: bytes, max_steps: int = 1 << 34,
                  expect: Optional[bytes] = None) -> Tuple[bytes, int]:
    mem = list(mem0)
    dev = Tape(tau)
    pc = 0
    steps = 0
    while pc != HALT_PC and steps < max_steps:
        A = mem[pc]
        B = mem[(pc + 1) & 0xFF]
        C = mem[(pc + 2) & 0xFF]
        r = (mem[B] - mem[A]) & 0xFF
        mem[B] = r
        pc = C if (r == 0 or r >= 0x80) else (pc + 3) & 0xFF
        n_before = len(dev.out)
        dev.apply(mem)
        if expect is not None and len(dev.out) > n_before:
            k = len(dev.out) - 1
            if k >= len(expect) or dev.out[k] != expect[k]:
                raise AssertionError(f"stream diverged at byte {k}")
        steps += 1
    return bytes(dev.out), steps

# =============================================================================
# The host netlist and its canonical serialization
# =============================================================================

STATE_LAYOUT = {"pc": [0, 8], "halt": [8, 1], "mem": [9, 256, 8]}
IO_CELLS = {"c_wr": C_WR, "c_eot": C_EOT, "c_rw": C_RW, "c_rd": C_RD,
            "r_in": R_IN, "r_out": R_OUT, "halt_pc": HALT_PC}
STATE_BITS = 8 + 1 + 2048


def host_netlist():
    """The clocked netlist of the host, assembled from its source description."""
    import host_netlist as H
    return H.build_subleq_step_net()


def sigma_host() -> bytes:
    """sigma(N_host): the canonical serialization of that netlist."""
    from netlist_io import sigma_of_net
    net, inputs, outputs = host_netlist()
    return sigma_of_net(net, inputs, outputs, "subleq8io", STATE_LAYOUT,
                        IO_CELLS)


def read_host() -> bytes:
    """The distributed serialization of the host."""
    return open(HOST_PATH, "rb").read()


def sha(b: bytes) -> str:
    return hashlib.sha256(b).hexdigest()


def tau_star(sigma: bytes, m: List[int]) -> bytes:
    return describe(field(sigma) + field(bytes(m)))


# =============================================================================
# Evaluators of the step map, each built from sigma alone
# =============================================================================

class _Transducer:
    """State marshalling shared by the threshold evaluators.

    A subclass supplies `step`, which maps a batch of state vectors to the next
    ones, and `n_in`, the width of the state. One step of the transducer of
    Definition 2.8 is that map followed by the device, so the run loop counts
    the step on which the halt bit is set and applies the device to it.
    """

    N = STATE_BITS
    _DEVCELLS = (C_WR, C_EOT, C_RW, C_RD, R_IN, R_OUT)

    def _vec(self, pc: int, mem: List[int]):
        torch = self.torch
        v = torch.zeros(self.N)
        for k in range(8):
            v[k] = (pc >> (7 - k)) & 1
        for j in range(256):
            for k in range(8):
                v[9 + j * 8 + k] = (mem[j] >> (7 - k)) & 1
        return v

    @staticmethod
    def _byte(vc, j: int) -> int:
        x = 0
        for k in range(8):
            x = (x << 1) | int(vc[9 + j * 8 + k])
        return x

    def _set_byte(self, v, j: int, val: int) -> None:
        for k in range(8):
            v[0, 9 + j * 8 + k] = float((val >> (7 - k)) & 1)

    def capture(self):
        """Capture one application of the map as a CUDA graph.

        The map is a fixed sequence of operations on fixed shapes, so the whole
        step replays as one graph launch in place of some hundreds. The captured
        map is compared against the eager one before it is used.
        """
        torch = self.torch
        assert self.device.startswith("cuda")
        self.gin = torch.zeros(1, self.N, device=self.device)
        # every buffer the body writes must stay alive for the life of the
        # graph, or the allocator will hand its memory to something else and
        # the replay will overwrite that instead
        self.gbuf = self._buffers()
        for _ in range(5):
            self._body(self.gin)
        torch.cuda.synchronize()
        self.graph = torch.cuda.CUDAGraph()
        with torch.cuda.graph(self.graph):
            self.gout = self._body(self.gin)
        torch.cuda.synchronize()
        gen = torch.Generator(device=self.device).manual_seed(3)
        probe = (torch.rand(1, self.N, generator=gen,
                            device=self.device) < 0.5).float()
        self.gin.copy_(probe)
        self.graph.replay()
        torch.cuda.synchronize()
        assert bool((self.gout == self._eager(probe)).all()), \
            "the captured graph differs from the eager step"
        return self

    def _replay(self, v):
        # the state lives in gin, which is ordinary memory; gout belongs to the
        # graph's private pool and is only ever read as a whole
        if v.data_ptr() != self.gin.data_ptr():
            self.gin.copy_(v)
        self.graph.replay()
        self.gin.copy_(self.gout)
        return self.gin

    def run(self, mem0: List[int], tau: bytes, max_steps: int,
            expect: Optional[bytes] = None, progress: int = 0,
            margin: bool = False) -> Tuple[bytes, int]:
        """Iterate the map with the device applied after each step.

        Only the halt bit and the six device cells cross to the host each step;
        the rest of the state stays on the accelerator. The six cells are read
        out in one operation and written back in one, so the cost of a step is
        the map and not the marshalling. With margin=True the minimum distance
        of any pre-activation from -1/2 along the whole trajectory is
        accumulated (dense mode only)."""
        import time
        torch = self.torch
        v = self._vec(0, mem0).unsqueeze(0).to(self.device)
        dev = Tape(tau)
        cells = list(self._DEVCELLS)
        cell_t = torch.tensor([9 + j * 8 + k for j in cells for k in range(8)],
                              device=self.device)
        pow2_t = torch.tensor([1 << (7 - k) for k in range(8)],
                              device=self.device, dtype=torch.float32)
        shift_t = torch.tensor([7 - k for k in range(8)])
        shadow = [0] * 256
        self.min_margin = float("inf")
        n = 0
        t0 = time.perf_counter()
        while n < max_steps:
            if margin:
                self._accumulate_margin(v)
            v = self.step(v)
            n += 1
            vals = torch.cat([v[0, 8:9],
                              (v[0, cell_t].reshape(len(cells), 8) * pow2_t)
                              .sum(-1)]).to("cpu").to(torch.int64).tolist()
            for c, j in enumerate(cells):
                shadow[j] = vals[1 + c]
            before = len(dev.out)
            dev.apply(shadow)
            new = torch.tensor([shadow[j] for j in cells], dtype=torch.int64)
            v[0, cell_t] = (((new.unsqueeze(-1) >> shift_t) & 1)
                            .reshape(-1).float().to(self.device))
            if expect is not None and len(dev.out) > before:
                k = len(dev.out) - 1
                if k >= len(expect) or dev.out[k] != expect[k]:
                    raise AssertionError(f"stream diverged at byte {k}")
            if vals[0] >= 1:
                break
            if progress and n % progress == 0:
                rate = n / (time.perf_counter() - t0)
                print(f"      {self.tag} {n:,} steps, {len(dev.out):,} bytes "
                      f"({rate:,.0f} steps/s)", flush=True)
        self.seconds = time.perf_counter() - t0
        return bytes(dev.out), n


class NetEvaluator(_Transducer):
    """The netlist of sigma, evaluated unit by unit.

    The units are grouped by depth and each reads its predecessors out of one
    signal vector, so no unit is evaluated before its predecessors and none is
    padded to a common width: the evaluation carries the netlist's own
    50,250 predecessor entries and nothing else.
    """

    tag = "net"

    def __init__(self, sigma: bytes, device: str = "cpu", graph: bool = False):
        import time
        import torch
        from netlist_io import net_of_sigma
        from reflect import Leveled
        t0 = time.perf_counter()
        self.torch = torch
        self.device = device
        net, inputs, outputs, meta = net_of_sigma(sigma)
        self.net, self.inputs, self.outputs, self.meta = net, inputs, outputs, meta
        assert len(inputs) == self.N, "state width disagrees with the layout"
        self.lev = Leveled(net, inputs, outputs, device=device)
        self.info = {"levels": len(self.lev.plan), "units": len(net.gates),
                     "entries": sum(len(i) for i, _ in net.gates.values())}
        self.graph = None
        if graph:
            self.capture()
        self.build_seconds = time.perf_counter() - t0

    def _buffers(self):
        return self.torch.zeros(self.lev.n_sig, 1, device=self.device)

    def _body(self, inp):
        V = self.gbuf
        lev = self.lev
        V.zero_()
        V[1] = 1.0
        V[lev.in_slots] = inp.T
        for idx, w, b, out in lev.plan:
            g = V[idx]
            V[out] = ((g * w[:, :, None]).sum(1) + b[:, None] >= 0).float()
        return V[lev.out_slots].T

    def _eager(self, v):
        return self.lev.step(v)

    def step(self, v):
        if self.graph is not None:
            return self._replay(v)
        return self.lev.step(v)


class LevEvaluator(_Transducer):
    """Lev(N) for the netlist N of sigma.

    `dense=True` materialises the matrices of Lemma 2.6 and iterates
    matrix-vector products; `dense=False` evaluates the same map from its
    nonzero entries, the identity rows included, which is the same function
    layer by layer. The two agree on every state tested by check_lev.
    """

    tag = "lev"

    def __init__(self, sigma: bytes, device: str = "cpu", dense: bool = True,
                 graph: bool = False):
        import time
        import torch
        from netlist_io import net_of_sigma
        from matrix8 import compile_net
        t0 = time.perf_counter()
        self.torch = torch
        self.device = device
        self.dense = dense
        net, inputs, outputs, meta = net_of_sigma(sigma)
        self.net, self.inputs, self.outputs, self.meta = net, inputs, outputs, meta
        assert len(inputs) == self.N, "state width disagrees with the layout"
        layers, info = compile_net(net, inputs, outputs)
        for W, _ in layers:
            assert set(torch.unique(W).tolist()) <= {-1, 0, 1}
        self.info = dict(info)
        self.info["size"] = sum(int(W.shape[0]) for W, _ in layers)
        self.info["nonzero"] = sum(int((W != 0).sum()) for W, _ in layers)
        if dense:
            self.W = [W.to(device=device, dtype=torch.float32)
                      for W, _ in layers]
            self.B = [b.to(device=device, dtype=torch.float32)
                      for _, b in layers]
        else:
            pad = device.startswith("cuda")
            self.plan = [self._sparsify(W, b, device, pad) for W, b in layers]
            self.widths = [int(W.shape[0]) for W, _ in layers]
        self.graph = None
        if graph:
            self.capture()
        self.build_seconds = time.perf_counter() - t0

    @staticmethod
    def _sparsify(W, b, device, pad):
        """One layer as the nonzero entries of its rows.

        With `pad`, every row of the layer is padded to the largest number of
        nonzero entries in it by an entry of weight zero, which contributes
        nothing to any pre-activation and makes the layer one gather and one
        reduction; that is the faster arrangement on the accelerator, where the
        cost is the number of operations issued. Without it the rows are grouped
        by their number of nonzero entries and no padding is read, which is the
        faster arrangement on a processor core.
        """
        import torch
        nz = (W != 0)
        counts = nz.sum(1)
        n = int(W.shape[0])
        sizes = [int(counts.max())] if pad else sorted(set(counts.tolist()))
        groups = []
        for k in sizes:
            rows = (torch.arange(n) if pad
                    else torch.nonzero(counts == k, as_tuple=False).flatten())
            idx = torch.zeros(len(rows), max(k, 1), dtype=torch.long)
            w = torch.zeros(len(rows), max(k, 1))
            for r, row in enumerate(rows.tolist()):
                cols = torch.nonzero(nz[row], as_tuple=False).flatten()
                idx[r, :len(cols)] = cols
                w[r, :len(cols)] = W[row, cols]
            groups.append((rows.to(device), idx.to(device),
                           w.to(device=device, dtype=torch.float32),
                           b[rows].to(device=device, dtype=torch.float32)))
        return groups

    def _buffers(self):
        import torch
        return [torch.zeros(1, n, device=self.device) for n in self.widths]

    def _sparse_step(self, v, buf):
        x = v
        for groups, y in zip(self.plan, buf):
            for rows, idx, w, b in groups:
                y[:, rows] = ((x[:, idx] * w).sum(-1) + b >= 0).float()
            x = y
        return x

    def _body(self, inp):
        return self._sparse_step(inp, self.gbuf)

    def _eager(self, v):
        return self._sparse_step(v, [self.torch.zeros(v.shape[0], n,
                                                      device=self.device)
                                     for n in self.widths])

    def step(self, v):
        if self.dense:
            for W, b in zip(self.W, self.B):
                v = ((v @ W.T + b) >= 0).float()
            return v
        if self.graph is not None:
            return self._replay(v)
        return self._eager(v)

    def _accumulate_margin(self, v):
        y = v
        for W, b in zip(self.W, self.B):
            pre = y @ W.T + b
            m = float((pre + 0.5).abs().min())
            if m < self.min_margin:
                self.min_margin = m
            y = (pre >= 0).float()

    def step_noisy(self, v, sigma: float, gen):
        """One step with additive Gaussian read noise per pre-activation and the
        comparator at -1/2.

        The two forms compute the same pre-activations, so the noise is drawn
        for the same quantities whichever is used; the sparse form draws it
        layer by layer over that layer's rows.
        """
        torch = self.torch
        if self.dense:
            for W, b in zip(self.W, self.B):
                pre = v @ W.T + b
                pre = pre + torch.randn(pre.shape, generator=gen,
                                        device=pre.device) * sigma
                v = (pre >= -0.5).float()
            return v
        x = v
        for groups, n in zip(self.plan, self.widths):
            y = torch.empty(x.shape[0], n, device=x.device)
            for rows, idx, w, b in groups:
                pre = (x[:, idx] * w).sum(-1) + b
                pre = pre + torch.randn(pre.shape, generator=gen,
                                        device=pre.device) * sigma
                y[:, rows] = (pre >= -0.5).float()
            x = y
        return x