coremlit 0.1.0

Safe, synchronous CoreML runtime for macOS (CPU/GPU/Neural Engine) with opt-in on-device multimodal pipelines: speech (Whisper STT, forced alignment, speaker diarization, Silero VAD), AudioSet sound-event tagging, and audio/text/image embeddings (CLAP, granite, SigLIP)
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
use super::*;
use crate::audio::whisper::options::DecodingOptions;

fn greedy(temperature: f32) -> GreedyTokenSampler {
  GreedyTokenSampler::new(temperature, 3, &DecodingOptions::new()).with_seed(42)
}

#[test]
fn argmax_at_zero_temperature_with_exact_logprob() {
  let logits = [1.0f32, 3.0, 2.0, 0.0];
  let result = greedy(0.0).sample(&logits);
  assert_eq!(result.token(), 1);
  assert!(!result.completed());
  // log softmax by hand: 3.0 - ln(e^1 + e^3 + e^2 + e^0)
  let log_z = logits.iter().map(|v| v.exp()).sum::<f32>().ln();
  assert!((result.logprob() - (3.0 - log_z)).abs() < 1e-5);
}

#[test]
fn eot_completes() {
  let logits = [0.0f32, 0.0, 0.0, 5.0]; // index 3 == eot
  let result = greedy(0.0).sample(&logits);
  assert_eq!(result.token(), 3);
  assert!(result.completed());
}

#[test]
fn nonzero_temperature_is_seed_deterministic_and_top_k_bounded() {
  // top_k = 5 (DecodingOptions default); only indices 0..5 by probability
  // can ever be drawn.
  let mut logits = vec![0.0f32; 16];
  for (i, v) in [9.0, 8.0, 7.0, 6.0, 5.0].iter().enumerate() {
    logits[i + 8] = *v; // the top-5 live at 8..13
  }
  let mut a = greedy(0.7);
  let mut b = greedy(0.7);
  for _ in 0..20 {
    let (ra, rb) = (a.sample(&logits), b.sample(&logits));
    assert_eq!(ra.token(), rb.token(), "same seed, same draw");
    assert!(
      (8..13).contains(&(ra.token() as usize)),
      "outside top-k drawn"
    );
    assert!(ra.logprob() <= 0.0);
  }
}

#[test]
fn fully_masked_logits_degenerate_without_panic_or_nan() {
  // Regression (task-4 review, Important): every entry masked (-inf)
  // panicked on the empty multinomial range at t != 0 and produced a NaN
  // logprob at t == 0. Both paths must return the same defined result.
  let masked = [f32::NEG_INFINITY; 8];
  for temperature in [0.0, 0.7] {
    let result = greedy(temperature).sample(&masked);
    assert_eq!(result.token(), 0, "t={temperature}");
    assert_eq!(result.logprob(), f32::NEG_INFINITY, "t={temperature}");
    assert!(!result.completed(), "t={temperature}");
  }
  // eot at the degenerate index still reports completion.
  let eot_zero = GreedyTokenSampler::new(0.7, 0, &DecodingOptions::new())
    .with_seed(42)
    .sample(&masked);
  assert!(eot_zero.completed());
}

#[test]
fn fully_masked_sample_does_not_consume_rng() {
  // The degenerate path must not consult the RNG: a sampler that first
  // saw a fully masked buffer draws the same stream afterwards as a
  // fresh same-seeded sampler.
  let logits: Vec<f32> = (0..16).map(|i| i as f32 * 0.25).collect();
  let mut interrupted = greedy(0.7);
  let mut fresh = greedy(0.7);
  interrupted.sample(&[f32::NEG_INFINITY; 16]);
  for _ in 0..10 {
    assert_eq!(
      interrupted.sample(&logits).token(),
      fresh.sample(&logits).token()
    );
  }
}

#[test]
fn drew_from_rng_tracks_real_rng_draws() {
  // F2 (codex round 3). The fallback ladder records reproducibility from THIS
  // fact, not from the temperature: it must be true iff `sample` actually
  // consulted the RNG -- a non-zero-temperature draw on a non-masked buffer.
  let logits = [1.0f32, 3.0, 2.0, 0.0];

  // Argmax (temperature 0) never draws.
  let mut argmax_sampler = greedy(0.0);
  assert!(
    !argmax_sampler.drew_from_rng(),
    "a fresh sampler has not drawn"
  );
  argmax_sampler.sample(&logits);
  assert!(
    !argmax_sampler.drew_from_rng(),
    "argmax decoding does not consult the RNG"
  );

  // A non-zero temperature draws, and the flag latches.
  let mut sampling = greedy(0.7);
  sampling.sample(&logits);
  assert!(
    sampling.drew_from_rng(),
    "a non-zero-temperature sample draws from the RNG"
  );

  // The all-masked degenerate path returns without drawing.
  let mut masked = greedy(0.7);
  masked.sample(&[f32::NEG_INFINITY; 4]);
  assert!(
    !masked.drew_from_rng(),
    "the all-masked degenerate path must not consult the RNG"
  );
}

#[test]
fn negative_temperature_wide_logits_sample_without_panic() {
  // F1 (codex round 3, High). Swift scales the logits by `1 / temperature`
  // FIRST, then softmaxes the *scaled* vector (`TokenSampler.swift:109-138`).
  // The port stabilized the softmax with `max(raw) * inv_t`, but for a
  // NEGATIVE temperature `inv_t < 0` reverses order, so that constant is the
  // *minimum* scaled value, not the max: the true-largest scaled entry then
  // computes `exp(scaled - min) = exp(huge) = +inf`, making `top_sum`
  // non-finite and panicking `random_range(0.0..inf)`.
  //
  // Reproduction (pre-fix): `sample([-10, 10])` at `-0.2` panicked. Post-fix
  // it must return a finite draw with no panic.
  let result = greedy(-0.2).sample(&[-10.0f32, 10.0]);
  assert!(result.logprob().is_finite(), "logprob must be finite");
  assert!((result.token() as usize) < 2, "token must index the logits");

  // Every drawn token must lie in the top-k of a numerically stable softmax
  // over the SCALED logits (`v / temperature`) -- exactly what Swift's
  // scale-then-softmax computes. Under a negative temperature the *smallest*
  // raw logit is the most probable, so a correct fix inverts the ordering
  // rather than overflowing.
  let wide = [-10.0f32, 10.0, -8.0, 5.0, -3.0, 2.0, 9.0, -1.0];
  let inv_t = 1.0f32 / -0.2;
  let scaled: Vec<f32> = wide.iter().map(|&v| v * inv_t).collect();
  let scaled_max = scaled.iter().copied().fold(f32::NEG_INFINITY, f32::max);
  let reference: Vec<f32> = scaled.iter().map(|&s| (s - scaled_max).exp()).collect();
  assert!(
    reference.iter().all(|p| p.is_finite()),
    "the reference stable softmax over scaled logits must stay finite"
  );
  // top_k defaults to 5: the five highest-probability scaled entries.
  let mut order: Vec<usize> = (0..wide.len()).collect();
  order.sort_by(|&a, &b| reference[b].total_cmp(&reference[a]));
  let top_k: std::collections::HashSet<usize> = order.into_iter().take(5).collect();

  let mut sampler = greedy(-0.2);
  for _ in 0..50 {
    let r = sampler.sample(&wide);
    assert!(
      r.logprob().is_finite(),
      "finite logprob under negative temperature"
    );
    assert!(
      top_k.contains(&(r.token() as usize)),
      "drawn token {} fell outside the stable-softmax top-k {top_k:?}",
      r.token()
    );
  }
}

#[test]
fn negative_temperature_masked_logits_sample_without_panic() {
  // F1 (codex round 4, High). The round-3 fix scaled every logit by `1/T`
  // and stabilized on `max(scaled)`, but a filter MASK (`-inf`, what
  // `decode::filter` writes for a suppressed token) is not a number to
  // scale: at a NEGATIVE temperature `-inf * (1/T)` flips to `+inf`, so the
  // masked index BECOMES the scaled max, the stabilizer subtraction yields
  // `+inf - +inf = NaN` across the vector, and `random_range(0.0..NaN)`
  // panics. The round-3 regression above never reached this regime -- it
  // used only finite logits, so no `-inf` was present to flip.
  //
  // Reproduction (pre-fix): `sample([0.0, NEG_INFINITY])` at `-0.2` panicked
  // on the NaN multinomial range. Post-fix: no panic, a finite log-prob, a
  // real RNG draw, and the masked index NEVER selected.
  let mut sampler = greedy(-0.2);
  for _ in 0..200 {
    let r = sampler.sample(&[0.0f32, f32::NEG_INFINITY]);
    assert!(
      r.logprob().is_finite(),
      "masked negative-temperature draw must have a finite log-prob"
    );
    assert_eq!(
      r.token(),
      0,
      "the masked index (1) must never be drawn -- a mask is a mask at any temperature sign"
    );
  }
  assert!(
    sampler.drew_from_rng(),
    "a non-zero-temperature draw on a non-masked-max buffer consults the RNG"
  );

  // The mask must stay excluded whatever the sign of the temperature, and
  // the finite entry must win regardless of where it sits. A wider mix of
  // finite and masked entries at a negative temperature: every draw lands on
  // a FINITE-logit index, never a masked one, and stays finite.
  let mixed = [
    f32::NEG_INFINITY,
    -4.0,
    f32::NEG_INFINITY,
    7.0,
    2.0,
    f32::NEG_INFINITY,
  ];
  let masked_indices = [0usize, 2, 5];
  let mut sampler = greedy(-0.35);
  for _ in 0..200 {
    let r = sampler.sample(&mixed);
    assert!(
      r.logprob().is_finite(),
      "finite log-prob under the mask + negative temperature"
    );
    assert!(
      !masked_indices.contains(&(r.token() as usize)),
      "a masked (-inf) index {} was drawn under negative temperature",
      r.token()
    );
  }
}

#[test]
fn tiny_temperature_overflow_sample_without_panic_or_nan() {
  // F1 sibling (codex round 4). A finite-but-tiny temperature drives `1/T`
  // (or a scaled logit) past the f32 range, so `v * (1/T)` overflows to
  // `+-inf` (or, for `v == 0` when `1/T` itself overflowed, the `0 * inf`
  // NaN). Either would break the stabilized softmax and violate `sample`'s
  // "panics only on empty logits" contract; the fix clamps the overflow to
  // the finite extremes and treats the `0 * inf` case as the true scaled
  // value `0`.
  for &temperature in &[1e-40f32, -1e-40, f32::MIN_POSITIVE, 1e-30] {
    let mut sampler = GreedyTokenSampler::new(temperature, 3, &DecodingOptions::new()).with_seed(1);
    for _ in 0..50 {
      let r = sampler.sample(&[0.0f32, 20.0, -20.0, 5.0]);
      assert!(
        r.logprob().is_finite() || r.logprob() == f32::NEG_INFINITY,
        "t={temperature}: log-prob must be finite or the degenerate -inf, never NaN"
      );
      assert!(
        (r.token() as usize) < 4,
        "t={temperature}: token must index the logits"
      );
    }
  }
}

#[test]
fn tiny_temperature_preserves_rank_and_logprob() {
  // F1 (codex round 14). At a subnormal temperature `1 / temperature` (or a
  // scaled logit) overflows f32; the pre-fix per-entry clamp then mapped every
  // same-sign finite logit to the SAME endpoint, so the stabilized softmax was
  // uniform over all of them. Concretely `[5.0, 20.0]` at `1e-40` clamped both
  // to `f32::MAX`, drew `[0.5, 0.5]`, and returned `ln(0.5)` for whichever token
  // it picked -- a possibly-wrong token, and a certainly-wrong logprob where the
  // limit is `0`. The overflow-only stable-difference path preserves order: the
  // temperature-sign-appropriate extreme takes probability 1 (logprob → 0) and
  // every other distinct finite logit collapses to probability 0 (logprob −∞),
  // independent of the seed. The existing
  // `tiny_temperature_overflow_sample_without_panic_or_nan` still pins the
  // no-panic / no-NaN contract this rank/logprob check sits on top of.
  //
  // Mutation: revert `sample` to the naive `(v * inv_t) - scaled_max` scaling
  // (drop the `saturates` stable path) and BOTH `±1e-40` blocks fail -- the draw
  // becomes uniform, so the token is seed-dependent and the logprob is `ln(1/n)`.
  let logits = [5.0f32, 20.0, -3.0, 12.0]; // distinct, finite; argmax = 1, argmin = 2
  for &(temperature, winner) in &[(1e-40f32, 1u32), (-1e-40f32, 2)] {
    for seed in 0..8u64 {
      // eot_token 3 is neither the argmax nor the argmin, so `completed` stays
      // false whichever extreme wins.
      let mut sampler =
        GreedyTokenSampler::new(temperature, 3, &DecodingOptions::new()).with_seed(seed);
      let result = sampler.sample(&logits);
      assert_eq!(
        result.token(),
        winner,
        "t={temperature}, seed={seed}: the order-preserving extreme must win, not a uniform draw"
      );
      assert!(
        result.logprob().abs() < 1e-6,
        "t={temperature}, seed={seed}: the point-mass winner's logprob → 0, got {}",
        result.logprob()
      );
      assert!(!result.completed(), "t={temperature}, seed={seed}");
    }
  }
}

#[test]
#[should_panic(expected = "non-empty logits")]
fn empty_logits_panic() {
  greedy(0.0).sample(&[]);
}

#[test]
fn finalize_appends_eot_once() {
  let sampler = greedy(0.0);
  let (mut tokens, mut logprobs) = (vec![1u32, 2], vec![-0.5f32, -0.25]);
  sampler.finalize(&mut tokens, &mut logprobs);
  assert_eq!(tokens, vec![1, 2, 3]);
  assert_eq!(logprobs, vec![-0.5, -0.25, 0.0]);
  sampler.finalize(&mut tokens, &mut logprobs); // idempotent: already ends in EOT
  assert_eq!(tokens.len(), 3);
}

// ---------------------------------------------------------------------
// derive_attempt_seed
// ---------------------------------------------------------------------

#[test]
fn derive_attempt_seed_is_pure_and_deterministic() {
  // Same tuple, same result -- every time, no hidden state.
  assert_eq!(
    derive_attempt_seed(1, 2, 3, 4),
    derive_attempt_seed(1, 2, 3, 4),
    "pure function: identical inputs must reproduce identical output"
  );
  // Each of the four coordinates independently changes the result: the
  // mixer folds every one into its own bijective `splitmix64` round, so
  // changing exactly one is *guaranteed* to change the output (see
  // `derive_attempt_seed`'s doc), not merely likely to.
  assert_ne!(
    derive_attempt_seed(1, 2, 3, 4),
    derive_attempt_seed(2, 2, 3, 4),
    "the base seed must change the derived seed"
  );
  assert_ne!(
    derive_attempt_seed(1, 2, 3, 4),
    derive_attempt_seed(1, 9, 3, 4),
    "worker_index must change the derived seed"
  );
  assert_ne!(
    derive_attempt_seed(1, 2, 3, 4),
    derive_attempt_seed(1, 2, 9, 4),
    "window_index must change the derived seed"
  );
  assert_ne!(
    derive_attempt_seed(1, 2, 3, 4),
    derive_attempt_seed(1, 2, 3, 9),
    "attempt_index must change the derived seed"
  );
}

#[test]
fn derive_attempt_seed_domain_separates_worker_and_window() {
  // Regression for the caller-side `offset + window_index` SUM alias
  // (coremlit#13): `transcribe_all` feeds each audio's global index as the
  // worker id while every task resets its window counter to 0, so
  // audio-0/window-1 (0 + 1) and audio-1/window-0 (1 + 0) summed to the
  // same coordinate `1` and shared one StdRng stream -- identical draws on
  // any shared logits shape (silent/repeated windows, or the MockBackend
  // that ignores encoder output). Passed as SEPARATE coordinates they must
  // derive different sub-seeds.
  for seed in [0u64, 1, 7, u64::MAX] {
    for attempt in 0..=5u64 {
      assert_ne!(
        derive_attempt_seed(seed, 0, 1, attempt),
        derive_attempt_seed(seed, 1, 0, attempt),
        "(worker=0, window=1) must not alias (worker=1, window=0) \
         [seed={seed} attempt={attempt}]"
      );
    }
  }
}

#[test]
fn derive_attempt_seed_has_no_zero_collapse_across_base_seeds() {
  // Regression for the mixer's XOR/zero alias (coremlit#13): the old
  // `splitmix64(seed ^ window)` mixer folded distinct coordinates together
  // because `splitmix64(0) == 0`. `(seed=0, window=0)` and
  // `(seed=1, window=1)` both derived 0, and `(seed=0, window=1)` aliased
  // `(seed=1, window=0)`. None of these may alias now, and the all-zero
  // tuple must not derive 0.
  assert_ne!(
    derive_attempt_seed(0, 0, 0, 0),
    0,
    "the all-zero tuple must not collapse to 0"
  );
  assert_ne!(
    derive_attempt_seed(0, 0, 0, 0),
    derive_attempt_seed(1, 0, 1, 0),
    "(seed=0, window=0) must not alias (seed=1, window=1)"
  );
  assert_ne!(
    derive_attempt_seed(0, 0, 1, 0),
    derive_attempt_seed(1, 0, 0, 0),
    "(seed=0, window=1) must not alias (seed=1, window=0)"
  );
}

#[test]
fn derive_attempt_seed_has_no_collisions_over_realistic_ranges() {
  // Statistical decorrelation check over a generous worker/window/attempt
  // range (temperature_fallback_count defaults to 5, so 0..=8 already
  // exceeds any default configuration; 0..64 windows covers long-form
  // audio's window count many times over; 0..64 workers covers a large
  // concurrent batch or VAD-chunk count). A broken derivation that ignored,
  // truncated, or summed any coordinate would collide immediately here.
  let mut seen = std::collections::HashSet::new();
  for worker in 0..64u64 {
    for window in 0..64u64 {
      for attempt in 0..=8u64 {
        let derived = derive_attempt_seed(0xABCD_1234_5678_9ABC, worker, window, attempt);
        assert!(
          seen.insert(derived),
          "collision at worker={worker} window={window} attempt={attempt}"
        );
      }
    }
  }
}

/// A seeded, non-degenerate (multi-candidate) sampler at `temperature =
/// 0.7` -- top_k defaults to 5, so this is a genuine multinomial draw
/// among several candidates, not a coin flip between two or a foregone
/// argmax.
fn seeded_sampler(seed: u64) -> GreedyTokenSampler {
  GreedyTokenSampler::new(0.7, 999, &DecodingOptions::new()).with_seed(seed)
}

fn draw_sequence(sampler: &mut GreedyTokenSampler, logits: &[f32], n: usize) -> Vec<u32> {
  (0..n).map(|_| sampler.sample(logits).token()).collect()
}

#[test]
fn attempt_seed_derivation_changes_sampled_draws_across_attempts() {
  // Proves the sub-seed derivation is actually wired into distinct SAMPLING
  // streams, not just distinct numbers in the abstract: two samplers seeded
  // from adjacent attempt indices at the same (worker, window) draw
  // different sequences from the exact same logits.
  //
  // Mutation check performed by hand (not left in the tree): temporarily
  // making `derive_attempt_seed` ignore `attempt_index` (returning the
  // same sub-seed for every attempt at a fixed window) made this
  // `assert_ne!` fail, confirming the test is sensitive to exactly the
  // bug class it exists to catch.
  let seed = 0xC0FFEE_u64;
  let worker = 2u64;
  let window = 3u64;
  let logits: Vec<f32> = (0..32).map(|i| i as f32 * 0.1 - 1.6).collect();

  let mut attempt0 = seeded_sampler(derive_attempt_seed(seed, worker, window, 0));
  let mut attempt1 = seeded_sampler(derive_attempt_seed(seed, worker, window, 1));
  let draws0 = draw_sequence(&mut attempt0, &logits, 20);
  let draws1 = draw_sequence(&mut attempt1, &logits, 20);
  assert_ne!(
    draws0, draws1,
    "different attempt_index must decorrelate the sampled stream"
  );

  // Reproducibility half: the identical tuple always replays the identical
  // stream (this is what makes a whole transcription reproducible from one
  // base seed).
  let mut replay0 = seeded_sampler(derive_attempt_seed(seed, worker, window, 0));
  let replay_draws0 = draw_sequence(&mut replay0, &logits, 20);
  assert_eq!(draws0, replay_draws0);
}

#[test]
fn attempt_seed_derivation_changes_sampled_draws_across_windows() {
  // Same proof as above, along the window_index coordinate: two different
  // windows at the same (worker, attempt) must not share a draw stream.
  let seed = 0xC0FFEE_u64;
  let worker = 2u64;
  let attempt = 1u64;
  let logits: Vec<f32> = (0..32).map(|i| i as f32 * 0.1 - 1.6).collect();

  let mut window0 = seeded_sampler(derive_attempt_seed(seed, worker, 0, attempt));
  let mut window1 = seeded_sampler(derive_attempt_seed(seed, worker, 1, attempt));
  let draws0 = draw_sequence(&mut window0, &logits, 20);
  let draws1 = draw_sequence(&mut window1, &logits, 20);
  assert_ne!(
    draws0, draws1,
    "different window_index must decorrelate the sampled stream"
  );
}

#[test]
fn attempt_seed_derivation_changes_sampled_draws_across_workers() {
  // The Class-A alias (coremlit#13) at the SAMPLED-STREAM level: the two
  // (worker, window) pairs the old `offset + window_index` sum collapsed --
  // (worker=0, window=1) and (worker=1, window=0) -- must now draw
  // different sequences from identical logits, exactly as two real
  // transcription windows sharing a logits shape would need.
  let seed = 0xC0FFEE_u64;
  let attempt = 0u64;
  let logits: Vec<f32> = (0..32).map(|i| i as f32 * 0.1 - 1.6).collect();

  let mut worker0_window1 = seeded_sampler(derive_attempt_seed(seed, 0, 1, attempt));
  let mut worker1_window0 = seeded_sampler(derive_attempt_seed(seed, 1, 0, attempt));
  let draws_a = draw_sequence(&mut worker0_window1, &logits, 20);
  let draws_b = draw_sequence(&mut worker1_window0, &logits, 20);
  assert_ne!(
    draws_a, draws_b,
    "(worker=0, window=1) and (worker=1, window=0) must not share a stream"
  );
}

// ---------------------------------------------------------------------
// argmax tie-break parity (H2, coremlit issue #41)
//
// Direct tests of the private `argmax` primitive against the pinned Swift
// oracle behavior (tests/whisper_swift_probes/probe_argmax2.out, macOS
// 26.5/M1 Max): `MLTensor.argmax` on the f32-cast logits (shipping macOS
// 15+ path) and BNNS `.argMax` (legacy) both return the FIRST index on
// every crafted tie and both skip NaN. Each test names its probe case.
// ---------------------------------------------------------------------

#[test]
fn argmax_tie_keeps_first_index() {
  // Probe cases `tie_2_and_5_small`, `adjacent_tie_100_101`,
  // `tie_0_and_last_small`, `vocab_manyway_tie_every_5000`: the first index
  // of an exact tie wins, whatever the tie's shape or vector size (the probe
  // confirmed size-independence n=16..51865).
  let mut distant = [0.0f32; 16];
  distant[2] = 5.0;
  distant[5] = 5.0;
  assert_eq!(
    argmax(&distant),
    2,
    "distant tie -> first (tie_2_and_5_small)"
  );

  let mut adjacent = [0.0f32; 1024];
  adjacent[100] = 5.0;
  adjacent[101] = 5.0;
  assert_eq!(
    argmax(&adjacent),
    100,
    "adjacent tie -> first (adjacent_tie_100_101)"
  );

  let mut zero_and_last = [0.0f32; 16];
  zero_and_last[0] = 5.0;
  zero_and_last[15] = 5.0;
  assert_eq!(
    argmax(&zero_and_last),
    0,
    "tie at 0 & last -> 0 (tie_0_and_last_small)"
  );

  let mut manyway = [0.0f32; 16];
  for v in manyway.iter_mut().step_by(5) {
    *v = 7.0;
  }
  assert_eq!(
    argmax(&manyway),
    0,
    "many-way tie -> first (vocab_manyway_tie_every_5000)"
  );
}

#[test]
fn argmax_all_equal_returns_zero() {
  // Probe cases `all_equal_zero_small`, `vocab_all_equal`: with no strict
  // maximum, index 0 wins.
  assert_eq!(argmax(&[0.0f32; 16]), 0);
  assert_eq!(argmax(&[1.0f32; 16]), 0);
}

#[test]
fn argmax_neg_infinity_floor_tie() {
  // Probe case `vocab_all_neginf_except_tie_123_50400`: every entry masked
  // to -inf except a tied finite pair -> the first of the pair.
  let mut v = [f32::NEG_INFINITY; 16];
  v[2] = 1.5;
  v[5] = 1.5;
  assert_eq!(
    argmax(&v),
    2,
    "first finite of the tied pair over an -inf floor"
  );
}

#[test]
fn argmax_signed_zero_ties_keep_first() {
  // Probe cases `signed_zero_tie_neg0_at_2_pos0_at_5` and
  // `signed_zero_tie_pos0_at_2_neg0_at_5`: IEEE `==` treats `-0.0 == +0.0`,
  // so the FIRST zero wins in BOTH orders (the two signed zeros are the max
  // over a -1 floor). Red-first discriminator: the previous `total_cmp`
  // body ranked `+0.0 > -0.0` and returned 5 for the first order.
  let mut neg_then_pos = [-1.0f32; 16];
  neg_then_pos[2] = -0.0;
  neg_then_pos[5] = 0.0;
  assert_eq!(argmax(&neg_then_pos), 2, "-0.0@2, +0.0@5 -> 2");

  let mut pos_then_neg = [-1.0f32; 16];
  pos_then_neg[2] = 0.0;
  pos_then_neg[5] = -0.0;
  assert_eq!(argmax(&pos_then_neg), 2, "+0.0@2, -0.0@5 -> 2");
}

#[test]
fn argmax_skips_nan_anywhere() {
  // Probe cases `nan_at_4_max_at_7` and `nan_at_0_max_at_7`: NaN is skipped
  // wherever it sits (including index 0, which must not seed `best`), and
  // the finite max wins. Red-first discriminator: the previous `total_cmp`
  // body ranked NaN above +inf and returned the NaN index (4, or 0 when the
  // NaN led).
  let mut nan_mid = [0.0f32; 16];
  nan_mid[4] = f32::NAN;
  nan_mid[7] = 5.0;
  assert_eq!(argmax(&nan_mid), 7, "NaN@4 skipped, max@7 wins");

  let mut nan_first = [0.0f32; 16];
  nan_first[0] = f32::NAN;
  nan_first[7] = 5.0;
  assert_eq!(
    argmax(&nan_first),
    7,
    "NaN@0 must not seed best, max@7 wins"
  );
}

#[test]
fn argmax_all_nan_pins_zero() {
  // Probe case `all_nan`: unspecified upstream (MLTensor -> 0, BNNS -> last);
  // this port pins 0, matching the shipping MLTensor path.
  assert_eq!(argmax(&[f32::NAN; 16]), 0);
}