Skip to content

Commit b53d7d2

Browse files
committed
Flux training: Implement scanned blocks, dynamic gradient checkpointing, and weight loading improvements
1 parent b2d31df commit b53d7d2

15 files changed

Lines changed: 580 additions & 252 deletions

File tree

src/maxdiffusion/__init__.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -367,7 +367,7 @@
367367
_import_structure["models.controlnet_flax"] = ["FlaxControlNetModel"]
368368
_import_structure["models.modeling_flax_utils"] = ["FlaxModelMixin"]
369369
_import_structure["models.unet_2d_condition_flax"] = ["FlaxUNet2DConditionModel"]
370-
_import_structure["models.flux.transformers.transformer_flux_flax"] = ["FluxTransformer2DModel"]
370+
_import_structure["models.flux.transformers.transformer_flux"] = ["FluxTransformer2DModel"]
371371
_import_structure["models.vae_flax"] = ["FlaxAutoencoderKL"]
372372
_import_structure["models.ltx_video.transformers.transformer3d"] = ["Transformer3DModel"]
373373
_import_structure["pipelines"].extend(["FlaxDiffusionPipeline"])
@@ -444,7 +444,7 @@
444444
from .models.controlnet_flax import FlaxControlNetModel
445445
from .models.modeling_flax_utils import FlaxModelMixin
446446
from .models.unet_2d_condition_flax import FlaxUNet2DConditionModel
447-
from .models.flux.transformers.transformer_flux_flax import FluxTransformer2DModel
447+
from .models.flux.transformers.transformer_flux import FluxTransformer2DModel
448448
from .models.ltx_video.transformers.transformer3d import Transformer3DModel
449449
from .models.vae_flax import FlaxAutoencoderKL
450450
from .pipelines import FlaxDiffusionPipeline

src/maxdiffusion/checkpointing/flux_checkpointer.py

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,7 @@
2727
FlaxAutoencoderKL,
2828
max_logging,
2929
)
30-
from maxdiffusion.models.flux.transformers.transformer_flux_flax import FluxTransformer2DModel
30+
from maxdiffusion.models.flux.transformers.transformer_flux import FluxTransformer2DModel
3131
from ..pipelines.flux.flux_pipeline import FluxPipeline
3232

3333
from transformers import (CLIPTokenizer, FlaxCLIPTextModel, FlaxT5EncoderModel, AutoTokenizer)
@@ -214,6 +214,11 @@ def load_diffusers_checkpoint(self):
214214
dtype=self.config.activations_dtype,
215215
weights_dtype=self.config.weights_dtype,
216216
precision=max_utils.get_precision(self.config),
217+
use_base2_exp=self.config.use_base2_exp,
218+
use_experimental_scheduler=self.config.use_experimental_scheduler,
219+
remat_policy=self.config.remat_policy,
220+
names_which_can_be_saved=self.config.names_which_can_be_saved,
221+
names_which_can_be_offloaded=self.config.names_which_can_be_offloaded,
217222
)
218223
transformer_eval_params = transformer.init_weights(
219224
rngs=self.rng, max_sequence_length=self.config.max_sequence_length, eval_only=True
@@ -279,6 +284,11 @@ def load_checkpoint(self, step=None, scheduler_class=None):
279284
weights_dtype=self.config.weights_dtype,
280285
precision=max_utils.get_precision(self.config),
281286
from_pt=self.config.from_pt,
287+
use_base2_exp=self.config.use_base2_exp,
288+
use_experimental_scheduler=self.config.use_experimental_scheduler,
289+
remat_policy=self.config.remat_policy,
290+
names_which_can_be_saved=self.config.names_which_can_be_saved,
291+
names_which_can_be_offloaded=self.config.names_which_can_be_offloaded,
282292
)
283293

284294
pipeline = FluxPipeline(

src/maxdiffusion/configs/base_flux_dev.yml

Lines changed: 33 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -63,6 +63,8 @@ jit_initializers: True
6363
from_pt: True
6464
split_head_dim: True
6565
attention: 'flash' # Supported attention: dot_product, flash, cudnn_flash_te
66+
use_base2_exp: False
67+
use_experimental_scheduler: False
6668
# If mask_padding_tokens is True, we pass in segment ids to splash attention to avoid attending to padding tokens.
6769
# Else we do not pass in segment ids and on vpu bound hardware like trillium this is faster.
6870
# However, when padding tokens are significant, this will lead to worse quality and should be set to True.
@@ -73,18 +75,18 @@ mask_padding_tokens: True
7375
# in cross attention q.
7476
attention_sharding_uniform: True
7577

76-
flash_block_sizes: {}
78+
#flash_block_sizes: {}
7779
# Use the following flash_block_sizes on v6e (Trillium) due to larger vmem.
78-
# flash_block_sizes: {
79-
# "block_q" : 1536,
80-
# "block_kv_compute" : 1536,
81-
# "block_kv" : 1536,
82-
# "block_q_dkv" : 1536,
83-
# "block_kv_dkv" : 1536,
84-
# "block_kv_dkv_compute" : 1536,
85-
# "block_q_dq" : 1536,
86-
# "block_kv_dq" : 1536
87-
# }
80+
flash_block_sizes: {
81+
"block_q" : 1536,
82+
"block_kv_compute" : 1536,
83+
"block_kv" : 1536,
84+
"block_q_dkv" : 1536,
85+
"block_kv_dkv" : 1536,
86+
"block_kv_dkv_compute" : 1536,
87+
"block_q_dq" : 1536,
88+
"block_kv_dq" : 1536
89+
}
8890
# GroupNorm groups
8991
norm_num_groups: 32
9092

@@ -147,9 +149,11 @@ mesh_axes: ['data', 'fsdp', 'context', 'tensor']
147149
# conv_in : conv.shape[2] weight
148150
# conv_out : conv.shape[-1] weight
149151
logical_axis_rules: [
150-
['batch', 'data'],
152+
['batch', ['data','fsdp']],
151153
['activation_batch', ['data','fsdp']],
152-
['activation_heads', 'tensor'],
154+
['activation_heads', 'fsdp'],
155+
['activation_length', 'context'],
156+
['activation_kv_length', 'context'],
153157
['activation_kv', 'tensor'],
154158
['mlp','tensor'],
155159
['embed','fsdp'],
@@ -188,7 +192,7 @@ dataset_type: 'tfrecord' # Options: 'tfrecord', 'hf', 'tf', 'grain', 'synthetic
188192
# 2. Optionally set synthetic_num_samples (null=infinite, or a number like 10000)
189193
# 3. Optionally override dimensions
190194
#
191-
# synthetic_num_samples: null # null for infinite, or set a number
195+
synthetic_num_samples: 1000 # null for infinite, or set a number
192196
#
193197
# Optional dimension overrides:
194198
# resolution: 512
@@ -218,6 +222,21 @@ transform_images_num_proc: 4
218222
reuse_example_batch: False
219223
enable_data_shuffling: True
220224

225+
# Defines the type of gradient checkpoint to enable.
226+
# NONE - means no gradient checkpoint
227+
# FULL - means full gradient checkpoint, whenever possible (minimum memory usage)
228+
# MATMUL_WITHOUT_BATCH - means gradient checkpoint for every linear/matmul operation,
229+
# except for ones that involve batch dimension - that means that all attention and projection
230+
# layers will have gradient checkpoint, but not the backward with respect to the parameters.
231+
# OFFLOAD_MATMUL_WITHOUT_BATCH - same as MATMUL_WITHOUT_BATCH but offload instead of recomputing.
232+
# CUSTOM - set names to offload and save.
233+
remat_policy: "FLUX_OPTIMIZED"
234+
# For CUSTOM policy set below, current annotations are for: attn_output, query_proj, key_proj, value_proj
235+
# xq_out, xk_out, ffn_activation
236+
names_which_can_be_saved: []
237+
names_which_can_be_offloaded: []
238+
flash_min_seq_length: 0
239+
221240
# checkpoint every number of samples, -1 means don't checkpoint.
222241
checkpoint_every: -1
223242
# enables one replica to read the ckpt then broadcast to the rest

src/maxdiffusion/generate_flux.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,7 @@
3232
from transformers import (CLIPTokenizer, FlaxCLIPTextModel, T5EncoderModel, FlaxT5EncoderModel, AutoTokenizer)
3333

3434
from maxdiffusion import FlaxAutoencoderKL, pyconfig, max_logging, max_utils
35-
from maxdiffusion.models.flux.transformers.transformer_flux_flax import FluxTransformer2DModel
35+
from maxdiffusion.models.flux.transformers.transformer_flux import FluxTransformer2DModel
3636
from maxdiffusion.train_utils import transformer_engine_context
3737
from maxdiffusion.max_utils import (
3838
device_put_replicated,
@@ -314,6 +314,9 @@ def run(config):
314314
dtype=config.activations_dtype,
315315
weights_dtype=config.weights_dtype,
316316
precision=get_precision(config),
317+
remat_policy=config.remat_policy,
318+
names_which_can_be_saved=config.names_which_can_be_saved,
319+
names_which_can_be_offloaded=config.names_which_can_be_offloaded,
317320
)
318321

319322
num_channels_latents = transformer.in_channels // 4

src/maxdiffusion/generate_flux_multi_res.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,7 @@
3131
from transformers import (CLIPTokenizer, FlaxCLIPTextModel, T5EncoderModel, FlaxT5EncoderModel, AutoTokenizer)
3232

3333
from maxdiffusion import FlaxAutoencoderKL, pyconfig, max_logging, max_utils
34-
from maxdiffusion.models.flux.transformers.transformer_flux_flax import FluxTransformer2DModel
34+
from maxdiffusion.models.flux.transformers.transformer_flux import FluxTransformer2DModel
3535
from maxdiffusion.max_utils import (
3636
device_put_replicated,
3737
get_memory_allocations,

src/maxdiffusion/models/__init__.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,7 +28,7 @@
2828
from .unet_2d_condition_flax import FlaxUNet2DConditionModel
2929
from .vae_flax import FlaxAutoencoderKL
3030
from .lora import *
31-
from .flux.transformers.transformer_flux_flax import FluxTransformer2DModel
31+
from .flux.transformers.transformer_flux import FluxTransformer2DModel
3232
from .ltx_video.transformers.transformer3d import Transformer3DModel
3333

3434
else:

src/maxdiffusion/models/attention_flax.py

Lines changed: 43 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -1824,6 +1824,8 @@ class FlaxFluxAttention(nn.Module):
18241824
out_axis_names: AxisNames = (BATCH, LENGTH, EMBED)
18251825
precision: jax.lax.Precision = None
18261826
qkv_bias: bool = False
1827+
use_base2_exp: bool = False
1828+
use_experimental_scheduler: bool = False
18271829

18281830
def setup(self):
18291831
if self.attention_kernel in {"flash", "cudnn_flash_te"} and self.mesh is None:
@@ -1843,6 +1845,8 @@ def setup(self):
18431845
flash_block_sizes=self.flash_block_sizes,
18441846
dtype=self.dtype,
18451847
float32_qk_product=False,
1848+
use_base2_exp=self.use_base2_exp,
1849+
use_experimental_scheduler=self.use_experimental_scheduler,
18461850
)
18471851

18481852
kernel_axes = ("embed", "heads")
@@ -1923,41 +1927,59 @@ def __call__(
19231927
attention_mask=None,
19241928
image_rotary_emb=None,
19251929
):
1926-
qkv_proj = self.qkv(hidden_states)
19271930
B, L = hidden_states.shape[:2]
1928-
H, D, K = self.heads, qkv_proj.shape[-1] // (self.heads * 3), 3
1929-
qkv_proj = qkv_proj.reshape(B, L, K, H, D).transpose(2, 0, 3, 1, 4)
1930-
query_proj, key_proj, value_proj = qkv_proj
1931+
# Deduce dimensions cleanly from class attributes
1932+
H, D = self.heads, self.dim_head
19311933

1932-
query_proj = self.query_norm(query_proj)
1934+
qkv_proj = self.qkv(hidden_states)
1935+
qkv_proj = checkpoint_name(qkv_proj, "img_qkv_proj")
1936+
1937+
qkv_proj = qkv_proj.reshape(B, L, 3, H, D)
1938+
query_proj, key_proj, value_proj = jnp.split(qkv_proj, 3, axis=2)
1939+
query_proj = query_proj.squeeze(2)
1940+
key_proj = key_proj.squeeze(2)
1941+
value_proj = value_proj.squeeze(2)
19331942

1943+
query_proj = self.query_norm(query_proj)
19341944
key_proj = self.key_norm(key_proj)
19351945

19361946
if encoder_hidden_states is not None:
1947+
B_enc, L_txt = encoder_hidden_states.shape[:2]
19371948
encoder_qkv_proj = self.encoder_qkv(encoder_hidden_states)
1938-
B, L = encoder_hidden_states.shape[:2]
1939-
H, D, K = self.heads, encoder_qkv_proj.shape[-1] // (self.heads * 3), 3
1940-
encoder_qkv_proj = encoder_qkv_proj.reshape(B, L, K, H, D).transpose(2, 0, 3, 1, 4)
1941-
encoder_query_proj, encoder_key_proj, encoder_value_proj = encoder_qkv_proj
1949+
encoder_qkv_proj = checkpoint_name(encoder_qkv_proj, "txt_qkv_proj")
1950+
encoder_qkv_proj = encoder_qkv_proj.reshape(B_enc, L_txt, 3, H, D)
1951+
enc_query_proj, enc_key_proj, enc_value_proj = jnp.split(encoder_qkv_proj, 3, axis=2)
1952+
enc_query_proj = enc_query_proj.squeeze(2)
1953+
enc_key_proj = enc_key_proj.squeeze(2)
1954+
enc_value_proj = enc_value_proj.squeeze(2)
19421955

1943-
encoder_query_proj = self.encoder_query_norm(encoder_query_proj)
1956+
encoder_query_proj = self.encoder_query_norm(enc_query_proj)
1957+
encoder_key_proj = self.encoder_key_norm(enc_key_proj)
19441958

1945-
encoder_key_proj = self.encoder_key_norm(encoder_key_proj)
1959+
query_proj = jnp.concatenate((encoder_query_proj, query_proj), axis=1)
1960+
key_proj = jnp.concatenate((encoder_key_proj, key_proj), axis=1)
1961+
value_proj = jnp.concatenate((enc_value_proj, value_proj), axis=1)
19461962

1947-
query_proj = jnp.concatenate((encoder_query_proj, query_proj), axis=2)
1948-
key_proj = jnp.concatenate((encoder_key_proj, key_proj), axis=2)
1949-
value_proj = jnp.concatenate((encoder_value_proj, value_proj), axis=2)
1950-
1951-
query_proj = nn.with_logical_constraint(query_proj, self.query_axis_names)
1952-
key_proj = nn.with_logical_constraint(key_proj, self.key_axis_names)
1953-
value_proj = nn.with_logical_constraint(value_proj, self.value_axis_names)
1963+
# query_proj = nn.with_logical_constraint(query_proj, self.query_axis_names)
1964+
# key_proj = nn.with_logical_constraint(key_proj, self.key_axis_names)
1965+
# value_proj = nn.with_logical_constraint(value_proj, self.value_axis_names)
19541966

19551967
image_rotary_emb = rearrange(image_rotary_emb, "n d (i j) -> n d i j", i=2, j=2)
1968+
1969+
query_proj = query_proj.swapaxes(1, 2)
1970+
key_proj = key_proj.swapaxes(1, 2)
19561971
query_proj, key_proj = apply_rope(query_proj, key_proj, image_rotary_emb)
1972+
query_proj = query_proj.swapaxes(1, 2)
1973+
key_proj = key_proj.swapaxes(1, 2)
1974+
1975+
query_proj = query_proj.reshape(B, -1, H * D)
1976+
key_proj = key_proj.reshape(B, -1, H * D)
1977+
value_proj = value_proj.reshape(B, -1, H * D)
19571978

1958-
query_proj = query_proj.transpose(0, 2, 1, 3).reshape(query_proj.shape[0], query_proj.shape[2], -1)
1959-
key_proj = key_proj.transpose(0, 2, 1, 3).reshape(key_proj.shape[0], key_proj.shape[2], -1)
1960-
value_proj = value_proj.transpose(0, 2, 1, 3).reshape(value_proj.shape[0], value_proj.shape[2], -1)
1979+
if encoder_hidden_states is not None:
1980+
query_proj = nn.with_logical_constraint(query_proj, self.query_axis_names)
1981+
key_proj = nn.with_logical_constraint(key_proj, self.key_axis_names)
1982+
value_proj = nn.with_logical_constraint(value_proj, self.value_axis_names)
19611983

19621984
attn_output = self.attention_op.apply_attention(query_proj, key_proj, value_proj, attention_mask=attention_mask)
19631985
context_attn_output = None

src/maxdiffusion/models/flux/__init__.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,4 +14,4 @@
1414
limitations under the License.
1515
"""
1616

17-
from .transformers.transformer_flux_flax import FluxTransformer2DModel
17+
from .transformers.transformer_flux import FluxTransformer2DModel

0 commit comments

Comments
 (0)