huggingface / huggingface/diffusers

Unstable flux transformer prediction

Abierto
#10,233 9 comentarios 1 reacción 0 asignados Ver en GitHub

Nadie ha tomado este issue todavía.

bug stale
Lenguaje dominante
Python
Estrellas
34.5k
Forks
7.3k
Merge medio
3 d 3 h
PR fusionados (30 d)
91

Descripción

Describe the bug

When feeding the same input (same noise, same prompt embed, same guidance scale and same timestep)to the flux.transformer, i got different output. I tried to fix the seed but nothing changed.

Reproduction
import torch
import math

from diffusers import FluxPipeline
from torch import Tensor

import random
import numpy as np

def setup_seed(seed=1234):
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    np.random.seed(seed)
    random.seed(seed)
    torch.backends.cudnn.deterministic = True
    torch.backends.cudnn.benchmark = False

def get_schedule(
    num_steps: int,
    image_seq_len: int,
    base_shift: float = 0.5,
    max_shift: float = 1.15,
    shift: bool = True,
) -> list:
    def get_lin_function(
        x1: float = 256, y1: float = 0.5, x2: float = 4096, y2: float = 1.15
    ):
        m = (y2 - y1) / (x2 - x1)
        b = y1 - m * x1
        return lambda x: m * x + b
    def time_shift(mu: float, sigma: float, t: Tensor):
        return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
    
    # extra step for zero
    timesteps = torch.linspace(1, 0, num_steps + 1, dtype=torch.float32)

    # shifting the schedule to favor high timesteps for higher signal images
    if shift:
        # estimate mu based on linear estimation between two points
        mu = get_lin_function(y1=base_shift, y2=max_shift)(image_seq_len)
        timesteps = time_shift(mu, 1.0, timesteps)

    return timesteps.tolist()

if __name__ == '__main__':
    setup_seed(1234) # fix the random seed

    DTYPE = torch.bfloat16
    device = 'cuda'
    pipe = FluxPipeline.from_pretrained('PATH/TO/MODEL', torch_dtype=DTYPE)
    pipe.to(device)

    with torch.no_grad():
        shape = (1, 16, 128, 128)
        rand_noise = torch.randn(shape, device=device, dtype=DTYPE)
        rand_noise = pipe._pack_latents(
            rand_noise,
            batch_size=shape[0],
            num_channels_latents=shape[1],
            height=shape[2],
            width=shape[3],
        )

        t_vec = torch.full((shape[0],), torch.tensor(1.), dtype=DTYPE, device=device)
        guidance_vec = torch.full((shape[0],), 2., device=device, dtype=DTYPE)
        latent_image_ids = pipe._prepare_latent_image_ids(
            shape[0], shape[2], shape[3], device, DTYPE,
        )

        prompt = 'A cat holding a sign that says hello world'
        prompt_embeds, pooled_prompt_embeds, text_ids = pipe.encode_prompt( # null text
            prompt=prompt, 
            prompt_2=prompt
        )
        while 1:
            flux_velocity = pipe.transformer(
                    hidden_states=rand_noise,
                    timestep=t_vec,
                    guidance=guidance_vec,
                    pooled_projections=pooled_prompt_embeds,
                    encoder_hidden_states=prompt_embeds,
                    txt_ids=text_ids,
                    img_ids=latent_image_ids,
                    joint_attention_kwargs=None,
                    return_dict=pipe,
                )[0]
            print(flux_velocity)
Logs
tensor([[[-1.7812,  0.5273, -0.5156,  ...,  0.6719,  0.5898,  0.5234],
         [-0.5664, -0.9258,  0.0598,  ...,  1.1719, -0.4688, -0.3555],
         [-1.2188, -0.2441, -0.4238,  ..., -0.3906, -0.2373,  1.0391],
         ...,
         [-0.0537, -0.2383, -0.0073,  ..., -1.3750, -1.1250,  0.8750],
         [-0.2812,  0.9062, -0.6250,  ..., -0.5195,  0.6797,  1.0391],
         [-0.7109, -0.5664, -0.7461,  ...,  0.2129, -0.5430,  0.6367]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5312, -0.5156,  ...,  0.6641,  0.5742,  0.5156],
         [-0.5703, -0.9297,  0.0569,  ...,  1.1719, -0.4766, -0.3594],
         [-1.2188, -0.2480, -0.4238,  ..., -0.3945, -0.2451,  1.0391],
         ...,
         [-0.0574, -0.2412, -0.0081,  ..., -1.3750, -1.1328,  0.8750],
         [-0.2793,  0.8984, -0.6289,  ..., -0.5234,  0.6719,  1.0391],
         [-0.7188, -0.5664, -0.7500,  ...,  0.2031, -0.5547,  0.6250]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7891,  0.5273, -0.5156,  ...,  0.6641,  0.5820,  0.5156],
         [-0.5625, -0.9258,  0.0588,  ...,  1.1719, -0.4668, -0.3594],
         [-1.2188, -0.2461, -0.4219,  ..., -0.3906, -0.2422,  1.0391],
         ...,
         [-0.0564, -0.2373, -0.0037,  ..., -1.3750, -1.1250,  0.8750],
         [-0.2793,  0.9062, -0.6211,  ..., -0.5234,  0.6758,  1.0391],
         [-0.7109, -0.5625, -0.7461,  ...,  0.2090, -0.5508,  0.6328]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7891,  0.5273, -0.5195,  ...,  0.6680,  0.5820,  0.5195],
         [-0.5664, -0.9297,  0.0598,  ...,  1.1719, -0.4688, -0.3594],
         [-1.2188, -0.2451, -0.4219,  ..., -0.3906, -0.2422,  1.0391],
         ...,
         [-0.0559, -0.2412, -0.0044,  ..., -1.3828, -1.1250,  0.8711],
         [-0.2793,  0.8984, -0.6328,  ..., -0.5273,  0.6719,  1.0391],
         [-0.7109, -0.5664, -0.7500,  ...,  0.2041, -0.5508,  0.6289]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5273, -0.5195,  ...,  0.6680,  0.5820,  0.5234],
         [-0.5625, -0.9219,  0.0603,  ...,  1.1797, -0.4688, -0.3594],
         [-1.2188, -0.2471, -0.4199,  ..., -0.3926, -0.2422,  1.0391],
         ...,
         [-0.0552, -0.2393, -0.0046,  ..., -1.3750, -1.1250,  0.8750],
         [-0.2734,  0.9062, -0.6250,  ..., -0.5234,  0.6758,  1.0391],
         [-0.7109, -0.5664, -0.7500,  ...,  0.2051, -0.5508,  0.6328]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7891,  0.5273, -0.5195,  ...,  0.6680,  0.5781,  0.5195],
         [-0.5664, -0.9258,  0.0581,  ...,  1.1797, -0.4707, -0.3594],
         [-1.2188, -0.2471, -0.4258,  ..., -0.3926, -0.2441,  1.0391],
         ...,
         [-0.0542, -0.2393, -0.0028,  ..., -1.3750, -1.1250,  0.8711],
         [-0.2812,  0.8984, -0.6328,  ..., -0.5234,  0.6758,  1.0312],
         [-0.7148, -0.5664, -0.7461,  ...,  0.2080, -0.5469,  0.6328]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5312, -0.5195,  ...,  0.6680,  0.5820,  0.5234],
         [-0.5625, -0.9258,  0.0593,  ...,  1.1797, -0.4688, -0.3574],
         [-1.2188, -0.2461, -0.4219,  ..., -0.3926, -0.2412,  1.0469],
         ...,
         [-0.0574, -0.2393, -0.0076,  ..., -1.3750, -1.1250,  0.8750],
         [-0.2793,  0.9023, -0.6289,  ..., -0.5234,  0.6758,  1.0391],
         [-0.7148, -0.5625, -0.7461,  ...,  0.2080, -0.5469,  0.6328]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5312, -0.5156,  ...,  0.6680,  0.5781,  0.5195],
         [-0.5664, -0.9258,  0.0588,  ...,  1.1797, -0.4707, -0.3594],
         [-1.2188, -0.2451, -0.4238,  ..., -0.3887, -0.2412,  1.0469],
         ...,
         [-0.0513, -0.2402, -0.0075,  ..., -1.3828, -1.1328,  0.8711],
         [-0.2773,  0.9023, -0.6250,  ..., -0.5234,  0.6758,  1.0312],
         [-0.7188, -0.5664, -0.7461,  ...,  0.2100, -0.5508,  0.6289]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5312, -0.5195,  ...,  0.6680,  0.5859,  0.5273],
         [-0.5625, -0.9258,  0.0549,  ...,  1.1797, -0.4668, -0.3555],
         [-1.2188, -0.2422, -0.4199,  ..., -0.3867, -0.2383,  1.0469],
         ...,
         [-0.0544, -0.2402, -0.0097,  ..., -1.3750, -1.1250,  0.8789],
         [-0.2754,  0.9023, -0.6289,  ..., -0.5234,  0.6797,  1.0391],
         [-0.7188, -0.5664, -0.7422,  ...,  0.2129, -0.5469,  0.6328]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5312, -0.5156,  ...,  0.6680,  0.5781,  0.5273],
         [-0.5664, -0.9258,  0.0598,  ...,  1.1797, -0.4668, -0.3574],
         [-1.2109, -0.2451, -0.4238,  ..., -0.3887, -0.2412,  1.0391],
         ...,
         [-0.0566, -0.2373, -0.0042,  ..., -1.3750, -1.1250,  0.8750],
         [-0.2773,  0.8984, -0.6250,  ..., -0.5234,  0.6797,  1.0391],
         [-0.7188, -0.5664, -0.7461,  ...,  0.2090, -0.5508,  0.6250]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5273, -0.5195,  ...,  0.6680,  0.5820,  0.5195],
         [-0.5664, -0.9297,  0.0605,  ...,  1.1719, -0.4688, -0.3594],
         [-1.2188, -0.2490, -0.4219,  ..., -0.3926, -0.2432,  1.0391],
         ...,
         [-0.0586, -0.2412, -0.0067,  ..., -1.3750, -1.1250,  0.8750],
         [-0.2773,  0.8984, -0.6289,  ..., -0.5234,  0.6719,  1.0391],
         [-0.7148, -0.5664, -0.7461,  ...,  0.2070, -0.5547,  0.6328]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5312, -0.5195,  ...,  0.6680,  0.5781,  0.5195],
         [-0.5664, -0.9258,  0.0608,  ...,  1.1719, -0.4707, -0.3594],
         [-1.2188, -0.2471, -0.4238,  ..., -0.3906, -0.2402,  1.0391],
         ...,
         [-0.0583, -0.2393, -0.0055,  ..., -1.3750, -1.1250,  0.8711],
         [-0.2773,  0.9023, -0.6250,  ..., -0.5195,  0.6758,  1.0391],
         [-0.7148, -0.5664, -0.7461,  ...,  0.2070, -0.5469,  0.6289]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5273, -0.5195,  ...,  0.6641,  0.5781,  0.5195],
         [-0.5664, -0.9297,  0.0579,  ...,  1.1797, -0.4727, -0.3594],
         [-1.2188, -0.2461, -0.4238,  ..., -0.3906, -0.2412,  1.0391],
         ...,
         [-0.0544, -0.2383, -0.0058,  ..., -1.3750, -1.1250,  0.8750],
         [-0.2793,  0.9062, -0.6250,  ..., -0.5195,  0.6758,  1.0391],
         [-0.7109, -0.5742, -0.7500,  ...,  0.2051, -0.5508,  0.6289]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5273, -0.5117,  ...,  0.6641,  0.5781,  0.5195],
         [-0.5664, -0.9297,  0.0579,  ...,  1.1719, -0.4727, -0.3633],
         [-1.2188, -0.2490, -0.4219,  ..., -0.3926, -0.2471,  1.0391],
         ...,
         [-0.0554, -0.2412, -0.0082,  ..., -1.3828, -1.1328,  0.8750],
         [-0.2793,  0.9023, -0.6250,  ..., -0.5234,  0.6719,  1.0312],
         [-0.7188, -0.5703, -0.7422,  ...,  0.2041, -0.5508,  0.6250]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5312, -0.5156,  ...,  0.6680,  0.5781,  0.5234],
         [-0.5664, -0.9258,  0.0576,  ...,  1.1719, -0.4707, -0.3613],
         [-1.2188, -0.2461, -0.4238,  ..., -0.3906, -0.2432,  1.0391],
         ...,
         [-0.0562, -0.2422, -0.0058,  ..., -1.3750, -1.1250,  0.8750],
         [-0.2773,  0.9023, -0.6289,  ..., -0.5273,  0.6719,  1.0391],
         [-0.7188, -0.5664, -0.7500,  ...,  0.2100, -0.5469,  0.6250]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5312, -0.5195,  ...,  0.6641,  0.5859,  0.5234],
         [-0.5664, -0.9258,  0.0593,  ...,  1.1719, -0.4688, -0.3574],
         [-1.2188, -0.2471, -0.4219,  ..., -0.3867, -0.2412,  1.0469],
         ...,
         [-0.0554, -0.2422, -0.0103,  ..., -1.3750, -1.1250,  0.8750],
         [-0.2793,  0.9062, -0.6289,  ..., -0.5234,  0.6758,  1.0312],
         [-0.7188, -0.5664, -0.7461,  ...,  0.2080, -0.5547,  0.6250]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812e+00,  5.3125e-01, -5.1562e-01,  ...,  6.6797e-01,
           5.8203e-01,  5.2344e-01],
         [-5.6641e-01, -9.2969e-01,  5.8594e-02,  ...,  1.1797e+00,
          -4.6680e-01, -3.5938e-01],
         [-1.2188e+00, -2.4414e-01, -4.2188e-01,  ..., -3.9062e-01,
          -2.3926e-01,  1.0469e+00],
         ...,
         [-5.6641e-02, -2.3730e-01, -9.3460e-04,  ..., -1.3750e+00,
          -1.1250e+00,  8.7500e-01],
         [-2.8125e-01,  9.0625e-01, -6.2891e-01,  ..., -5.1953e-01,
           6.7578e-01,  1.0469e+00],
         [-7.0703e-01, -5.6250e-01, -7.4609e-01,  ...,  2.1094e-01,
          -5.4688e-01,  6.3672e-01]]], device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7891,  0.5273, -0.5156,  ...,  0.6641,  0.5820,  0.5156],
         [-0.5664, -0.9297,  0.0559,  ...,  1.1719, -0.4707, -0.3633],
         [-1.2188, -0.2461, -0.4219,  ..., -0.3926, -0.2402,  1.0391],
         ...,
         [-0.0535, -0.2393, -0.0095,  ..., -1.3828, -1.1250,  0.8750],
         [-0.2773,  0.9023, -0.6250,  ..., -0.5195,  0.6758,  1.0391],
         [-0.7109, -0.5703, -0.7500,  ...,  0.2061, -0.5508,  0.6328]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5312, -0.5156,  ...,  0.6641,  0.5781,  0.5234],
         [-0.5625, -0.9258,  0.0569,  ...,  1.1719, -0.4727, -0.3613],
         [-1.2188, -0.2500, -0.4219,  ..., -0.3926, -0.2451,  1.0391],
         ...,
         [-0.0569, -0.2412, -0.0060,  ..., -1.3828, -1.1250,  0.8711],
         [-0.2812,  0.8984, -0.6250,  ..., -0.5234,  0.6719,  1.0312],
         [-0.7188, -0.5664, -0.7500,  ...,  0.1992, -0.5586,  0.6289]]],
       device='cuda:3', dtype=torch.bfloat16)
tensor([[[-1.7812,  0.5312, -0.5156,  ...,  0.6641,  0.5742,  0.5234],
         [-0.5664, -0.9258,  0.0596,  ...,  1.1797, -0.4746, -0.3633],
         [-1.2188, -0.2480, -0.4238,  ..., -0.3926, -0.2432,  1.0391],
         ...,
         [-0.0576, -0.2393, -0.0052,  ..., -1.3828, -1.1328,  0.8750],
         [-0.2773,  0.8984, -0.6328,  ..., -0.5312,  0.6680,  1.0312],
         [-0.7188, -0.5703, -0.7500,  ...,  0.2002, -0.5586,  0.6289]]],
       device='cuda:3', dtype=torch.bfloat16)
System Info
  • 🤗 Diffusers version: 0.31.0
  • Platform: Linux-5.4.0-144-generic-x86_64-with-glibc2.35
  • Running on Google Colab?: No
  • Python version: 3.10.14
  • PyTorch version (GPU?): 2.0.0+cu117 (True)
  • Flax version (CPU?/GPU?/TPU?): not installed (NA)
  • Jax version: not installed
  • JaxLib version: not installed
  • Huggingface_hub version: 0.25.2
  • Transformers version: 4.45.2
  • Accelerate version: 1.0.1
  • PEFT version: not installed
  • Bitsandbytes version: not installed
  • Safetensors version: 0.4.5
  • xFormers version: not installed
  • Accelerator: NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
  • Using GPU in script?: Yes
  • Using distributed or parallel set-up in script?: No
Who can help?

@DN6 @yiyixuxu @sayakpaul

Guía de contribución

Abrir la guía de contribución

Primeros pasos

  1. Lee el issue completo y luego la guía de contribución del proyecto.
  2. Comenta en el issue que vas a ocuparte — evita que dos personas hagan lo mismo.
  3. Haz un fork del repositorio y trabaja en una rama.
  4. Abre un pull request que haga referencia al número del issue.

Línea de trabajo

Comienza ejecutando la reproducción proporcionada con FluxPipeline y la llamada repetida a pipe.transformer en CUDA usando bfloat16. Compara las salidas repetidas mientras compruebas la ruta del transformer y la configuración de determinismo de PyTorch; el trabajo estará terminado cuando se identifique el origen de la variación y se establezcan predicciones repetidas deterministas, con una prueba de regresión si el proyecto tiene una ubicación de prueba aplicable.

Escrito por el modelo de indexación a partir del texto del issue.

Evaluación

Stack tecnológico
python, pytorch
Área
machine-learning
Tipo de issue
Error
Dificultad
4/5
Tiempo estimado
3-5 días
Estado de actividad
Estancado
Claridad
Necesita aclaración
Aptitud para principiantes
25/100

Recibe los nuevos issues en tu correo

Un resumen breve de issues de GitHub para principiantes.