modelscope / modelscope/DiffSynth-Studio

RuntimeError: FlashAttention only support fp16 and bf16 data type

Open
#628 5 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
13.1k
Forks
1.3k
Avg merge
13h 12m
Merged PRs (30d)
45

Description

Traceback (most recent call last): -- File "/root/code/train.py", line 134, in handle trainer.fit(model=model, datamodule=datamodule) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/trainer/trainer.py", line 561, in fit call._call_and_handle_interrupt( File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/trainer/call.py", line 48, in _call_and_handle_interrupt return trainer_fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/trainer/trainer.py", line 599, in _fit_impl self._run(model, ckpt_path=ckpt_path) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/trainer/trainer.py", line 1012, in _run results = self._run_stage() File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/trainer/trainer.py", line 1056, in _run_stage self.fit_loop.run() File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/loops/fit_loop.py", line 216, in run self.advance() File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/loops/fit_loop.py", line 455, in advance self.epoch_loop.run(self._data_fetcher) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/loops/training_epoch_loop.py", line 150, in run self.advance(data_fetcher) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/loops/training_epoch_loop.py", line 320, in advance batch_output = self.automatic_optimization.run(trainer.optimizers[0], batch_idx, kwargs) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/loops/optimization/automatic.py", line 192, in run self._optimizer_step(batch_idx, closure) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/loops/optimization/automatic.py", line 270, in _optimizer_step call._call_lightning_module_hook( File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/trainer/call.py", line 176, in _call_lightning_module_hook output = fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/core/module.py", line 1302, in optimizer_step optimizer.step(closure=optimizer_closure) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/core/optimizer.py", line 154, in step step_output = self._strategy.optimizer_step(self._optimizer, closure, **kwargs) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/strategies/strategy.py", line 239, in optimizer_step return self.precision_plugin.optimizer_step(optimizer, model=model, closure=closure, **kwargs) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/plugins/precision/amp.py", line 76, in optimizer_step return super().optimizer_step(optimizer, model=model, closure=closure, **kwargs) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/plugins/precision/precision.py", line 123, in optimizer_step return optimizer.step(closure=closure, **kwargs) File "/opt/conda/lib/python3.10/site-packages/torch/optim/optimizer.py", line 487, in wrapper out = func(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/torch/optim/optimizer.py", line 91, in _use_grad ret = func(self, *args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/torch/optim/adamw.py", line 197, in step loss = closure() File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/plugins/precision/precision.py", line 109, in _wrap_closure closure_result = closure() File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/loops/optimization/automatic.py", line 146, in __call__ self._result = self.closure(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/torch/utils/_contextlib.py", line 116, in decorate_context return func(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/loops/optimization/automatic.py", line 131, in closure step_output = self._step_fn() File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/loops/optimization/automatic.py", line 319, in _training_step training_step_output = call._call_strategy_hook(trainer, "training_step", *kwargs.values()) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/trainer/call.py", line 328, in _call_strategy_hook output = fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/pytorch_lightning/strategies/strategy.py", line 391, in training_step return self.lightning_module.training_step(*args, **kwargs) File "/root/code/src/trainers/custom_trainer.py", line 252, in training_step noise_pred = self.pipe.denoising_model()( File "/opt/conda/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl return self._call_impl(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl return forward_call(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/diffsynth/models/wan_video_dit.py", line 362, in forward x = torch.utils.checkpoint.checkpoint( File "/opt/conda/lib/python3.10/site-packages/torch/_compile.py", line 32, in inner return disable_fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/torch/_dynamo/eval_frame.py", line 632, in _fn return fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/torch/utils/checkpoint.py", line 496, in checkpoint ret = function(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/diffsynth/models/wan_video_dit.py", line 355, in custom_forward return module(*inputs) File "/opt/conda/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl return self._call_impl(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl return forward_call(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/diffsynth/models/wan_video_dit.py", line 218, in forward x = self.gate(x, gate_msa, self.self_attn(input_x, freqs)) File "/opt/conda/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl return self._call_impl(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl return forward_call(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/diffsynth/models/wan_video_dit.py", line 145, in forward x = self.attn(q, k, v) File "/opt/conda/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl return self._call_impl(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl return forward_call(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/diffsynth/models/wan_video_dit.py", line 119, in forward x = flash_attention(q=q, k=k, v=v, num_heads=self.num_heads) File "/opt/conda/lib/python3.10/site-packages/diffsynth/models/wan_video_dit.py", line 46, in flash_attention x = flash_attn.flash_attn_func(q, k, v) File "/opt/conda/lib/python3.10/site-packages/flash_attn/flash_attn_interface.py", line 880, in flash_attn_func return FlashAttnFunc.apply( File "/opt/conda/lib/python3.10/site-packages/torch/autograd/function.py", line 575, in apply return super().apply(*args, **kwargs)  # type: ignore[misc] File "/opt/conda/lib/python3.10/site-packages/flash_attn/flash_attn_interface.py", line 546, in forward out, q, k, v, out_padded, softmax_lse, S_dmask, rng_state = _flash_attn_forward( File "/opt/conda/lib/python3.10/site-packages/flash_attn/flash_attn_interface.py", line 52, in _flash_attn_forward out, q, k, v, out_padded, softmax_lse, S_dmask, rng_state = flash_attn_cuda.fwd( RuntimeError: FlashAttention only support fp16 and bf16 data type

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start with /root/code/src/trainers/custom_trainer.py:252 and the attention path in diffsynth/models/wan_video_dit.py, especially flash_attention at line 46. Reproduce the training run from /root/code/train.py:134 and inspect the tensor types at the failing call. Done means the same training path completes without the reported RuntimeError.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
30/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.