Skip to content

Commit be317b5

Browse files
committed
print statement fixes
1 parent 3740aee commit be317b5

7 files changed

Lines changed: 47 additions & 18 deletions

File tree

src/maxdiffusion/generate_flux2klein.py

Lines changed: 20 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
from absl import app
2121
import jax
2222
import jax.numpy as jnp
23+
import numpy as np
2324
import flax
2425
from flax import linen as nn
2526
from flax.linen import partitioning as nn_partitioning
@@ -458,8 +459,7 @@ def unbox_fn(x):
458459
print("\n" + "=" * 80)
459460
print("🚀 Running initial dry run (Warmup Pass) to compile XLA graphs...")
460461
print("=" * 80)
461-
t_warmup_start = time.time()
462-
pipeline(
462+
_, warmup_trace = pipeline(
463463
prompt=active_prompts,
464464
params=params,
465465
vae_params=vae_params,
@@ -478,13 +478,16 @@ def unbox_fn(x):
478478
output_dir=config.output_dir,
479479
output_name="flux2klein_warmup.png",
480480
)
481-
warmup_time = time.time() - t_warmup_start
481+
warmup_time = (
482+
warmup_trace.get("prompt_encoding", 0.0)
483+
+ warmup_trace.get("denoise_loop", 0.0)
484+
+ warmup_trace.get("vae_decode", 0.0)
485+
)
482486

483487
print("\n" + "=" * 80)
484488
print("⏱️ Running timed pass at full TPU speed...")
485489
print("=" * 80)
486-
t_main_start = time.time()
487-
pipeline(
490+
_, main_trace = pipeline(
488491
prompt=active_prompts,
489492
params=params,
490493
vae_params=vae_params,
@@ -503,14 +506,22 @@ def unbox_fn(x):
503506
output_dir=config.output_dir,
504507
output_name="flux2klein_generated_image.png",
505508
)
506-
main_time = time.time() - t_main_start
509+
main_time = (
510+
main_trace.get("prompt_encoding", 0.0) + main_trace.get("denoise_loop", 0.0) + main_trace.get("vae_decode", 0.0)
511+
)
507512

508513
print("\n" + "=" * 80)
509-
print("📊 FLUX.2-KLEIN LATENCY & TIMING BREAKDOWN")
514+
print("📊 FLUX.2-KLEIN LATENCY & TIMING BREAKDOWN (PURE MODEL INFERENCE)")
510515
print("=" * 80)
511-
print(f"1) Total Model Loading & Placement Time: {load_time:.2f} seconds ⏱️")
516+
print(f"1) Total Model Loading & Placement Time: {load_time:.2f} seconds ⏱️")
512517
print(f"2) Cold-Start / Warmup Pass (XLA Compilation): {warmup_time:.2f} seconds ⏱️")
513-
print(f"3) Main Warmed-Up Pass (Full-Speed Inference): {main_time:.2f} seconds ⏱️")
518+
print(f" - Qwen3 Encoding: {warmup_trace.get('prompt_encoding', 0.0):.2f}s")
519+
print(f" - Flux Denoising: {warmup_trace.get('denoise_loop', 0.0):.2f}s")
520+
print(f" - VAE Decoding: {warmup_trace.get('vae_decode', 0.0):.2f}s")
521+
print(f"3) Main Warmed-Up Pass (Pure Model Inference): {main_time:.2f} seconds ⏱️")
522+
print(f" - Qwen3 Encoding: {main_trace.get('prompt_encoding', 0.0):.2f}s")
523+
print(f" - Flux Denoising: {main_trace.get('denoise_loop', 0.0):.2f}s")
524+
print(f" - VAE Decoding: {main_trace.get('vae_decode', 0.0):.2f}s")
514525
print("=" * 80)
515526

516527
print("\n=======================================================")

src/maxdiffusion/models/flux/util.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -454,6 +454,7 @@ def load_and_convert_vae_weights(safetensors_path, jax_params):
454454
"""Loads PyTorch VAE weights from safetensors, maps them to JAX, and extracts BN stats."""
455455
from safetensors.torch import load_file
456456
import torch
457+
import numpy as np
457458
import flax
458459
import jax.numpy as jnp
459460

src/maxdiffusion/models/qwen3_flax.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313
# limitations under the License.
1414

1515
import math
16-
from typing import Any, List, Optional, Tuple
16+
from typing import Any, List, Optional, Tuple, Union
1717
import flax.linen as nn
1818
import jax
1919
import jax.numpy as jnp

src/maxdiffusion/pipelines/flux/flux2klein_pipeline.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,12 +15,13 @@
1515
import gc
1616
import os
1717
import time
18-
from typing import List, Union
18+
from typing import List, Union, Optional
1919
from PIL import Image
2020

2121
import jax
2222
import jax.numpy as jnp
2323
import numpy as np
24+
import flax
2425
from flax.linen import partitioning as nn_partitioning
2526

2627
from ..pipeline_flax_utils import FlaxDiffusionPipeline

src/maxdiffusion/tests/flux2klein/generate_flux2klein_test.py

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -81,12 +81,16 @@ def prepare_text_ids(batch_size, seq_len):
8181
class GenerateFlux2KleinTest(unittest.TestCase):
8282

8383
def test_generate_random_latents_shape(self):
84+
import torch
85+
8486
latents = torch.randn((2, 32, 1024 // 8, 512 // 8))
8587
expected_shape = (2, 32, 1024 // 8, 512 // 8)
8688
self.assertEqual(tuple(latents.shape), expected_shape)
8789

8890
def test_load_golden_latents_shape(self):
8991
# Deterministically generate initial latents on CPU using seed 0
92+
import torch
93+
9094
generator = torch.Generator(device="cpu").manual_seed(0)
9195
latents_pt = torch.randn((1, 32, 64, 64), generator=generator, dtype=torch.float32)
9296
expected_shape = (1, 32, 512 // 8, 512 // 8)
@@ -110,6 +114,7 @@ def test_qwen3_prompt_embeddings(self):
110114
self.fail(f"Failed to generate prompt embeddings: {e}")
111115

112116
def test_context_embedder_projection(self):
117+
import os
113118
import torch
114119
import jax
115120
import jax.numpy as jnp
@@ -329,11 +334,14 @@ def test_scheduler_timesteps_parity(self):
329334

330335
def test_attention_blocks_parity(self):
331336
"""Verifies that JAX joint-attention (double) and single-stream blocks match PyTorch golden outputs."""
337+
import os
338+
import torch
332339
import jax
333340
import jax.numpy as jnp
334341
import flax
335342
from flax.linen import partitioning as nn_partitioning
336343
from jax.sharding import Mesh
344+
from safetensors.torch import load_file
337345

338346
from maxdiffusion.models.flux.transformers.transformer_flux_flax import FluxTransformer2DModel
339347
from maxdiffusion import pyconfig
@@ -491,6 +499,7 @@ def test_attention_blocks_parity(self):
491499

492500
def test_full_transformer_and_multistep_parity(self):
493501
"""Verifies full JAX transformer forward pass (all blocks) and 4-step denoising loop parity against PyTorch."""
502+
import os
494503
import torch
495504
import jax
496505

@@ -704,6 +713,7 @@ def callback_fn(pipe, step_idx, timestep, callback_kwargs):
704713

705714
def test_vae_decoder_parity(self):
706715
"""Verifies JAX FlaxAutoencoderKL VAE Decoder parity against PyTorch."""
716+
import os
707717
import jax
708718
import jax.numpy as jnp
709719
import numpy as np
@@ -876,7 +886,7 @@ def get_w(key):
876886

877887
# 7. Compare raw decoder output
878888
diff_raw = jnp.abs(jax_decoder_out.sample - golden_decoder_out_pt)
879-
print("\n[VAE DIAG] Raw Decoder Output Comparison:")
889+
print(f"\n[VAE DIAG] Raw Decoder Output Comparison:")
880890
print(f"[VAE DIAG] Max absolute diff: {jnp.max(diff_raw)}")
881891
print(f"[VAE DIAG] Mean absolute diff: {jnp.mean(diff_raw)}")
882892

@@ -904,6 +914,7 @@ def test_10_point_isolated_parity_benchmark(self):
904914

905915
def test_swiglu_mlp_math_parity(self):
906916
"""Verifies SwiGLU MLP mathematical parity between JAX and PyTorch."""
917+
import torch
907918
import jax
908919
import jax.numpy as jnp
909920
import numpy as np
@@ -956,6 +967,7 @@ def test_attention_math_parity(self):
956967
pyconfig._config.keys["activations_dtype"] = "float32"
957968
config = pyconfig.config
958969

970+
import torch
959971
import jax
960972
import jax.numpy as jnp
961973
import numpy as np
@@ -1052,6 +1064,8 @@ def test_attention_math_parity(self):
10521064

10531065
def test_flowmatch_scheduler_parity(self):
10541066
"""Verifies FlowMatch Euler Discrete Scheduler stepping parity against PyTorch."""
1067+
import torch
1068+
import jax
10551069
import jax.numpy as jnp
10561070
import numpy as np
10571071
from diffusers import FlowMatchEulerDiscreteScheduler
@@ -1103,6 +1117,7 @@ def test_double_transformer_block_dummy_parity(self):
11031117
pyconfig._config.keys["activations_dtype"] = "float32"
11041118
config = pyconfig.config
11051119

1120+
import torch
11061121
import jax
11071122
import jax.numpy as jnp
11081123
import numpy as np
@@ -1231,6 +1246,7 @@ def test_single_transformer_block_dummy_parity(self):
12311246
pyconfig._config.keys["activations_dtype"] = "float32"
12321247
config = pyconfig.config
12331248

1249+
import torch
12341250
import jax
12351251
import jax.numpy as jnp
12361252
import numpy as np

src/maxdiffusion/tests/flux2klein/test_4b_e2e_parity.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -496,19 +496,19 @@ def jitted_vae_decode(v_params, latents_unpatched):
496496
l2_direct = np.sqrt(np.sum((jax_bf16_np - pt_bf16_np) ** 2))
497497
max_err_direct = np.max(np.abs(jax_bf16_np - pt_bf16_np))
498498

499-
print("\n--- Leg 2 (PyTorch CPU BF16) vs Leg 1 (PyTorch CPU FP32) Baseline ---")
499+
print(f"\n--- Leg 2 (PyTorch CPU BF16) vs Leg 1 (PyTorch CPU FP32) Baseline ---")
500500
print(f" SSIM: {ssim_pt_bf16:.6f}")
501501
print(f" RMSE: {rmse_pt_bf16:.4f} / 255")
502502
print(f" L2 Distance: {l2_pt_bf16:.4f}")
503503
print(f" Max Absolute Err: {max_err_pt_bf16:.4f} / 255")
504504

505-
print("\n--- Leg 3 (JAX TPU BF16) vs Leg 1 (PyTorch CPU FP32) Parity ---")
505+
print(f"\n--- Leg 3 (JAX TPU BF16) vs Leg 1 (PyTorch CPU FP32) Parity ---")
506506
print(f" SSIM: {ssim_jax_bf16:.6f} (Target: > 0.88)")
507507
print(f" RMSE: {rmse_jax_bf16:.4f} / 255")
508508
print(f" L2 Distance: {l2_jax_bf16:.4f}")
509509
print(f" Max Absolute Err: {max_err_jax_bf16:.4f} / 255")
510510

511-
print("\n--- Leg 3 (JAX TPU BF16) vs Leg 2 (PyTorch CPU BF16) Direct Alignment ---")
511+
print(f"\n--- Leg 3 (JAX TPU BF16) vs Leg 2 (PyTorch CPU BF16) Direct Alignment ---")
512512
print(f" SSIM: {ssim_direct:.6f}")
513513
print(f" RMSE: {rmse_direct:.4f} / 255")
514514
print(f" L2 Distance: {l2_direct:.4f}")

src/maxdiffusion/tests/flux2klein/test_9b_e2e_parity.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -506,19 +506,19 @@ def jitted_vae_decode(v_params, latents_unpatched):
506506
l2_direct = np.sqrt(np.sum((jax_bf16_np - pt_bf16_np) ** 2))
507507
max_err_direct = np.max(np.abs(jax_bf16_np - pt_bf16_np))
508508

509-
print("\n--- Leg 2 (PyTorch CPU BF16) vs Leg 1 (PyTorch CPU FP32) Baseline ---")
509+
print(f"\n--- Leg 2 (PyTorch CPU BF16) vs Leg 1 (PyTorch CPU FP32) Baseline ---")
510510
print(f" SSIM: {ssim_pt_bf16:.6f}")
511511
print(f" RMSE: {rmse_pt_bf16:.4f} / 255")
512512
print(f" L2 Distance: {l2_pt_bf16:.4f}")
513513
print(f" Max Absolute Err: {max_err_pt_bf16:.4f} / 255")
514514

515-
print("\n--- Leg 3 (JAX TPU BF16) vs Leg 1 (PyTorch CPU FP32) Parity ---")
515+
print(f"\n--- Leg 3 (JAX TPU BF16) vs Leg 1 (PyTorch CPU FP32) Parity ---")
516516
print(f" SSIM: {ssim_jax_bf16:.6f} (Target: > 0.88)")
517517
print(f" RMSE: {rmse_jax_bf16:.4f} / 255")
518518
print(f" L2 Distance: {l2_jax_bf16:.4f}")
519519
print(f" Max Absolute Err: {max_err_jax_bf16:.4f} / 255")
520520

521-
print("\n--- Leg 3 (JAX TPU BF16) vs Leg 2 (PyTorch CPU BF16) Direct Alignment ---")
521+
print(f"\n--- Leg 3 (JAX TPU BF16) vs Leg 2 (PyTorch CPU BF16) Direct Alignment ---")
522522
print(f" SSIM: {ssim_direct:.6f}")
523523
print(f" RMSE: {rmse_direct:.4f} / 255")
524524
print(f" L2 Distance: {l2_direct:.4f}")

0 commit comments

Comments
 (0)