feat(flux2klein): Faster Loading + Compliation - #456
Conversation
There was a problem hiding this comment.
Code Review
This pull request introduces several performance optimizations and robustness improvements to the FLUX.2-Klein pipeline, including concurrent Ahead-of-Time (AOT) compilation of XLA graphs, fused VAE decoding, direct target-dtype weight loading, and nearest-neighbor upsampling using JAX broadcasting. It also adds support for multiple inference repetitions, profiling, and fallback configurations when local files are missing. The review feedback highlights two critical issues: first, removing the return statement in partition_prompts causes the function to return None when prompt truncation is triggered; second, both the try and except blocks for snapshot downloading use local_files_only=True, which will fail completely if the model is not already cached locally instead of falling back to an online download.
1369c52 to
2bd0d37
Compare
2bd0d37 to
d39315d
Compare
acf2281 to
c23dc32
Compare
…6D broadcast upsampling, and dynamic topology sharding preservation
c23dc32 to
56a1119
Compare
| # PHASE A: Encode Prompt (Qwen3) | ||
| # --------------------------------------------------------------------- | ||
| if prompts is None: | ||
| prompts = ["A dog running in a field with butterflies and tall grass"] |
There was a problem hiding this comment.
let's not hardcode the prompt here, we can read it from a yaml file (same to other prompts that are hard coded)
| pt_config = AutoConfig.from_pretrained(text_encoder_path, local_files_only=True) | ||
| except Exception: | ||
| depth_val = getattr(config, "depth", 24) | ||
| hf_repo = "black-forest-labs/FLUX.2-klein-9B" if depth_val in (24, -1) else "black-forest-labs/FLUX.2-klein-4B" |
| os.makedirs(config.output_dir, exist_ok=True) | ||
|
|
||
| if hasattr(config, "per_device_batch_size") and config.per_device_batch_size > 0: | ||
| calculated_batch_size = int(config.per_device_batch_size * jax.device_count()) |
There was a problem hiding this comment.
maybe add a check/assertion here to avoid calculated_batch_size=0 case, like if per_device_batch_size=0.25 and device count = 2, we will have int(0.25 * 2) == 0
There was a problem hiding this comment.
why did the ref image change? I did a run on the two commits and got ssim(old_ref, new_ref, channel_axis=-1, data_range=255) -> 0.7981
There was a problem hiding this comment.
The ref image we had in the initial commit for the Flux 2 Klein model was stale (the implementation from the first commit would only give a 0.85 SSIM against it). The original reference image was generated using FP32 for the Qwen3 text encoder, whereas the committed code on maxdiffusion would only use bfloat16 (hence causing 1->0.85 SSIM drop).
There was also an issue with how batch sizes were being handled in the initial commit, where the smoke test would always generate multiple images. This commit fixes that issue and generates only a single image. The SSIM dropped from (0.85->0.798) from this change. This drop was because XLA will fuse operations differently when we use a different batch size, and hence cause slight rounding errors in bfloat16 to accumulate and lead to a slightly lower SSIM.
None of these changes affect the model implementation, so the outputs are still bit-wise identical.
On a v6-4, we see the following speedups (BS=8, fsdp) for the 9B model:
On a v7-8:
-Warm model loading: 540s -> 140s (for cold loads, use a Hyperdisk mount)
-XLA compilation/warmup-pass: 82s -> 49s
-Inference unchanged
So, on the v6-4: We bring the loading+compilation time down to <30s. On a v7-8, its <4 min. The main bottleneck is the hardware constraint of how fast it can load the model.
Small tweak to smoke test threshold (0.8 to 0.78), all other tests passing.