From d951a313b1bcc8a2eb48b43f8222c672e982848e Mon Sep 17 00:00:00 2001 From: sukru tikves Date: Tue, 15 Sep 2026 09:37:05 -0700 Subject: [PATCH] Switch SD3 and Flux2 noise generation to TorchRandomSource SD3 and Flux2 generated their initial latent noise with the NumPy RNG, so they did not match torch.randn, which the diffusers reference pipelines use. Three fixes to the torch RNG path: 1. TorchRandomSource.normalArray re-filled the trailing partial block with nextDouble(); PyTorch uses nextFloat(). PyTorch also fills every element with uniforms before applying Box-Muller, then re-fills the last 16 for the remainder. Matching this corrects count=17 and every size that is not a multiple of 16. 2. generateNoise(.torch) used the scalar nextNormal() loop, which yields a different sequence than torch.randn's batch-16 fill. Route it through normalArray. 3. SD3Pipeline and Flux2Pipeline called generateNoise without a sourceType and defaulted to .numPy. Set both to .torch. Adds TorchRandomSource parity tests against torch.randn reference values (seeds 0/42; counts 16/17/32/4096) covering the scalar, batch, boundary, remainder, and realistic-shape cases. Closes #151 --- .../Pipelines/Flux2Pipeline.swift | 2 +- .../Pipelines/SD3Pipeline.swift | 3 +- .../RNG/NoiseGeneration.swift | 4 +- .../RNG/TorchRandomSource.swift | 35 +++--- .../SchedulerTests.swift | 115 ++++++++++++++++++ 5 files changed, 142 insertions(+), 17 deletions(-) diff --git a/swift/Sources/CoreAIDiffusionPipeline/Pipelines/Flux2Pipeline.swift b/swift/Sources/CoreAIDiffusionPipeline/Pipelines/Flux2Pipeline.swift index d292bbd8..150626f0 100644 --- a/swift/Sources/CoreAIDiffusionPipeline/Pipelines/Flux2Pipeline.swift +++ b/swift/Sources/CoreAIDiffusionPipeline/Pipelines/Flux2Pipeline.swift @@ -224,7 +224,7 @@ public struct Flux2Pipeline: DiffusionPipeline { // 4. Generate noise [1, inChannels, spatialSide, spatialSide] let latentShape = [1, inChannels, spatialSide, spatialSide] let latentCount = latentShape.reduce(1, *) - let noise = generateNoise(count: latentCount, seed: configuration.seed) + let noise = generateNoise(count: latentCount, seed: configuration.seed, sourceType: .torch) let noisePacked = packLatentsSpatialFlatten( noise, channels: inChannels, height: spatialSide, width: spatialSide) diff --git a/swift/Sources/CoreAIDiffusionPipeline/Pipelines/SD3Pipeline.swift b/swift/Sources/CoreAIDiffusionPipeline/Pipelines/SD3Pipeline.swift index 2a5cd6fa..37b32d12 100644 --- a/swift/Sources/CoreAIDiffusionPipeline/Pipelines/SD3Pipeline.swift +++ b/swift/Sources/CoreAIDiffusionPipeline/Pipelines/SD3Pipeline.swift @@ -109,7 +109,8 @@ public struct SD3Pipeline: DiffusionPipeline { // 3. Initial noise var latents = generateNoise( count: latentShape.reduce(1, *), - seed: configuration.seed) + seed: configuration.seed, + sourceType: .torch) // 4. Scheduler (SD3 flow matching; plain shift, no dynamic mu) let scheduler = DiscreteFlowScheduler( diff --git a/swift/Sources/CoreAIDiffusionPipeline/RNG/NoiseGeneration.swift b/swift/Sources/CoreAIDiffusionPipeline/RNG/NoiseGeneration.swift index 1bda1b83..75f0b439 100644 --- a/swift/Sources/CoreAIDiffusionPipeline/RNG/NoiseGeneration.swift +++ b/swift/Sources/CoreAIDiffusionPipeline/RNG/NoiseGeneration.swift @@ -11,6 +11,8 @@ public enum RandomSourceType: Sendable { } /// Generate Gaussian noise (mean 0, stdev 1) using the specified random source. +/// `.torch` routes through `normalArray` to match `torch.randn` (batch-16 fill), +/// not the scalar `nextNormal()` loop. public func generateNoise(count: Int, seed: UInt32, sourceType: RandomSourceType = .numPy) -> [Float] { switch sourceType { case .numPy: @@ -18,7 +20,7 @@ public func generateNoise(count: Int, seed: UInt32, sourceType: RandomSourceType return (0..= 16 elements. + /// Matches `torch.randn(shape, dtype=.float)`, including PyTorch's `normal_fill_16` + /// batch-16 Box-Muller path for counts >= 16 (see the step comments below). public mutating func normalArray(_ shape: [Int], mean: Double = 0.0, stdev: Double = 1.0) -> [Float] { let count = shape.reduce(1, *) guard count >= 16 else { return (0..)`. + + // Reference: torch.manual_seed(42); [torch.randn(1, dtype=torch.float64).item() for _ in range(8)] + @Test("Torch scalar path matches Python reference (seed=42)") + func torchScalarParity() { + var rng = TorchRandomSource(seed: 42) + let expected: [Double] = [ + 0.3366903544, 0.1288094051, 0.2344623634, 0.2303330279, + -1.1228563767, -0.1863282993, 2.2082013356, -0.6379970568, + ] + for (i, exp) in expected.enumerated() { + let got = rng.nextNormal(mean: 0, stdev: 1) + #expect(abs(got - exp) < 1e-6, "scalar[\(i)]: got \(got), expected \(exp)") + } + } + + // Reference: torch.manual_seed(42); torch.randn(32, dtype=torch.float32) + @Test("Torch batch path matches Python torch.randn (seed=42, count=32)") + func torchBatchParity32() { + var rng = TorchRandomSource(seed: 42) + let expected: [Float] = [ + 1.9269150496, 1.4872841835, 0.9007171988, -2.1055214405, + 0.6784184575, -1.2345449924, -0.0430674814, -1.6046669483, + -0.7521361709, 1.6487228870, -0.3924786448, -1.4036067724, + -0.7278812528, -0.5594298840, -0.7688389421, 0.7624453902, + 1.6423169374, -0.1595973223, -0.4973974824, 0.4395892322, + -0.7581311464, 1.0783176422, 0.8008005023, 1.6806205511, + 1.2791243792, 1.2964228392, 0.6104664803, 1.3347377777, + -0.2316243201, 0.0417594910, -0.2515752614, 0.8598585129, + ] + let got = rng.normalArray([32], mean: 0, stdev: 1) + for (i, (g, e)) in zip(got, expected).enumerated() { + #expect(abs(g - e) < 1e-4, "batch[\(i)]: got \(g), expected \(e)") + } + } + + // Reference: torch.manual_seed(42); torch.randn(16, dtype=torch.float32) + @Test("Torch batch path boundary (exactly 16 elements)") + func torchBatchBoundary16() { + var rng = TorchRandomSource(seed: 42) + let expected: [Float] = [ + 1.9269150496, 1.4872841835, 0.9007171988, -2.1055214405, + 0.6784184575, -1.2345449924, -0.0430674814, -1.6046669483, + -0.7521361709, 1.6487228870, -0.3924786448, -1.4036067724, + -0.7278812528, -0.5594298840, -0.7688389421, 0.7624453902, + ] + let got = rng.normalArray([16], mean: 0, stdev: 1) + for (i, (g, e)) in zip(got, expected).enumerated() { + #expect(abs(g - e) < 1e-4, "boundary[\(i)]: got \(g), expected \(e)") + } + } + + // Reference: torch.manual_seed(42); torch.randn(17, dtype=torch.float32) + @Test("Torch batch path remainder (17 elements)") + func torchBatchRemainder17() { + var rng = TorchRandomSource(seed: 42) + let expected: [Float] = [ + 1.9269150496, -0.1595973223, -0.4973974824, 0.4395892322, + -0.7581311464, 1.0783176422, 0.8008005023, 1.6806205511, + 0.3558597863, 1.2964228392, 0.6104664803, 1.3347377777, + -0.2316243201, 0.0417594910, -0.2515752614, 0.8598585129, + -0.3097269237, + ] + let got = rng.normalArray([17], mean: 0, stdev: 1) + for (i, (g, e)) in zip(got, expected).enumerated() { + #expect(abs(g - e) < 1e-4, "remainder[\(i)]: got \(g), expected \(e)") + } + } + + // Reference: torch.manual_seed(0); torch.randn(32, dtype=torch.float32) + @Test("Torch batch path matches Python torch.randn (seed=0, count=32)") + func torchBatchParity0() { + var rng = TorchRandomSource(seed: 0) + let expected: [Float] = [ + -1.1258398294, -1.1523602009, -0.2505785823, -0.4338788390, + 0.8487103581, 0.6920092106, -0.3160127699, -2.1152195930, + 0.3222749233, -1.2633347511, 0.3499831855, 0.3081339002, + 0.1198415086, 1.2376579046, 1.1167771816, -0.2472776473, + -1.3526537418, -1.6959313154, 0.5666505098, 0.7935084105, + 0.5988394618, -1.5550950766, -0.3413603008, 1.8530061245, + 0.7501894236, -0.5854971409, -0.1733970195, 0.1834779233, + 1.3893661499, 1.5863343477, 0.9462983608, -0.8436768055, + ] + let got = rng.normalArray([32], mean: 0, stdev: 1) + for (i, (g, e)) in zip(got, expected).enumerated() { + #expect(abs(g - e) < 1e-4, "batch[\(i)]: got \(g), expected \(e)") + } + } + + // Reference: torch.manual_seed(42); t = torch.randn([1,16,16,16], dtype=torch.float32) + @Test("Torch batch path realistic shape (4096 elements, seed=42)") + func torchBatchRealisticShape() { + var rng = TorchRandomSource(seed: 42) + let got = rng.normalArray([1, 16, 16, 16], mean: 0, stdev: 1) + #expect(got.count == 4096) + + let expectedFirst: [Float] = [ + 1.9269150496, 1.4872841835, 0.9007171988, -2.1055214405, + 0.6784184575, -1.2345449924, -0.0430674814, -1.6046669483, + ] + let expectedLast: [Float] = [ + 1.5869791508, 0.1421326697, 0.3760589659, -0.7916260362, + 2.6677629948, -0.1403129250, 0.9416193962, -0.0118428767, + ] + for (i, (g, e)) in zip(got.prefix(8), expectedFirst).enumerated() { + #expect(abs(g - e) < 1e-4, "first[\(i)]: got \(g), expected \(e)") + } + for (i, (g, e)) in zip(got.suffix(8), expectedLast).enumerated() { + #expect(abs(g - e) < 1e-4, "last[\(i)]: got \(g), expected \(e)") + } + } } @Suite("Schedulers")