Skip to content

How to generate white background images instead of colorful? #14814

Description

@emadyounan

Hi, I'm trying to train stable-diffusion-v1-5/stable-diffusion-v1-5 with custom data with LoraConfig to generate a type of 2D asset.
So I have a dataset of transparent PNG images, and I added to my prompt text for every image a white background after the caption
And in my script, I added BG_COLOR = (255, 255, 255)
and that's my code

BASE_MODEL_ID = "stable-diffusion-v1-5/stable-diffusion-v1-5"
DATASET_DIR = r"Dataset/10_postapo_icon_style"
OUTPUT_DIR = r"output_lora_test"
OUTPUT_NAME = "postapo_icon_lora_test.safetensors"

RESOLUTION = 512  # لو حصل OOM قلّلها لـ 448 أو 384
BATCH_SIZE = 1  # على كارت أحدث ممكن 2-4 (بيقلّل عدد الخطوات في الـ epoch، فظبط الـ epochs/LR حسبه)
LEARNING_RATE = 1e-4
NUM_EPOCHS = 15
LORA_RANK = 16
SAVE_EVERY = 200  # checkpoint كل كام خطوة
# None = تلقائي (بيشتغل لو الـ VRAM أقل من 12 GB). True/False لو عايز تفرضه
GRADIENT_CHECKPOINTING = None

# نفس لون الخلفية اللي اتكتبت بيه الكابشنات (BG_COLOR في auto_caption.py)
BG_COLOR = (255, 255, 255)
FALLBACK_CAPTION = "postapo_style, 2d, white background"  # لو صورة ملهاش .txt
IMAGE_EXTS = (".png", ".jpg", ".jpeg", ".webp")
# =========================================================

DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
USE_AMP = DEVICE == "cuda"

VRAM_GB = 0.0
if USE_AMP:
    props = torch.cuda.get_device_properties(0)
    VRAM_GB = props.total_memory / 1024**3
    # bfloat16 لكروت Ampere وما فوق (RTX 30/40/50)، وfp16 للأقدم.
    # بنعتمد على compute capability لأن is_bf16_supported() بترجع True غلط على كروت قديمة
    AMP_DTYPE = torch.bfloat16 if props.major >= 8 else torch.float16
    print(f"GPU: {props.name} | {VRAM_GB:.1f} GB | compute capability {props.major}.{props.minor}")
else:
    AMP_DTYPE = torch.float32

# الأجزاء المجمّدة بنفس الـ dtype، وأوزان الـ LoRA نفسها fp32
WEIGHT_DTYPE = AMP_DTYPE
USE_GRAD_SCALER = AMP_DTYPE == torch.float16  # bf16 مش محتاج scaler
if GRADIENT_CHECKPOINTING is None:
    GRADIENT_CHECKPOINTING = 0 < VRAM_GB < 12

print(
    f"--- Starting LoRA Training on {DEVICE} ({WEIGHT_DTYPE}), "
    f"gradient checkpointing={GRADIENT_CHECKPOINTING} ---"
)


def load_square(path, size: int) -> Image.Image:
    """يلصق الشفافية على BG_COLOR، وبعدين يعمل padding لمربع بدل ما يمط الصورة."""
    img = Image.open(path)
    if img.mode in ("RGBA", "LA") or (img.mode == "P" and "transparency" in img.info):
        img = img.convert("RGBA")
        base = Image.new("RGBA", img.size, BG_COLOR + (255,))
        base.alpha_composite(img)
        img = base.convert("RGB")
    else:
        img = img.convert("RGB")

    w, h = img.size
    side = max(w, h)
    canvas = Image.new("RGB", (side, side), BG_COLOR)
    canvas.paste(img, ((side - w) // 2, (side - h) // 2))
    return canvas.resize((size, size), Image.LANCZOS)


class TextImageDataset(Dataset):
    def __init__(self, folder_path, resolution=512):
        folder = Path(folder_path)
        self.image_paths = sorted(
            p for p in folder.iterdir() if p.suffix.lower() in IMAGE_EXTS
        )
        self.resolution = resolution
        self.transform = transforms.Compose(
            [transforms.ToTensor(), transforms.Normalize([0.5], [0.5])]
        )

    def __len__(self):
        return len(self.image_paths)

    def __getitem__(self, idx):
        img_path = self.image_paths[idx]
        pixel_values = self.transform(load_square(img_path, self.resolution))

        txt_path = img_path.with_suffix(".txt")
        if txt_path.exists():
            caption = txt_path.read_text(encoding="utf-8").strip()
        else:
            caption = FALLBACK_CAPTION

        return {"pixel_values": pixel_values, "caption": caption}


# 3. تحميل القطع الأساسية
tokenizer = CLIPTokenizer.from_pretrained(BASE_MODEL_ID, subfolder="tokenizer")
text_encoder = CLIPTextModel.from_pretrained(
    BASE_MODEL_ID, subfolder="text_encoder"
).to(DEVICE, dtype=WEIGHT_DTYPE)
vae = AutoencoderKL.from_pretrained(BASE_MODEL_ID, subfolder="vae").to(
    DEVICE, dtype=WEIGHT_DTYPE
)
unet = UNet2DConditionModel.from_pretrained(BASE_MODEL_ID, subfolder="unet").to(
    DEVICE, dtype=WEIGHT_DTYPE
)
noise_scheduler = DDPMScheduler.from_pretrained(BASE_MODEL_ID, subfolder="scheduler")

for m in (unet, vae, text_encoder):
    m.requires_grad_(False)
text_encoder.eval()
vae.eval()

# 4. LoRA (بالطريقة الرسمية في diffusers: add_adapter بدل get_peft_model)
lora_config = LoraConfig(
    r=LORA_RANK,
    lora_alpha=LORA_RANK,
    target_modules=["to_q", "to_k", "to_v", "to_out.0"],
    lora_dropout=0.0,
    bias="none",
)
unet.add_adapter(lora_config)

# أوزان الـ LoRA لازم تفضل fp32 عشان التدريب يبقى مستقر
lora_params = [p for p in unet.parameters() if p.requires_grad]
for p in lora_params:
    p.data = p.data.float()
print(f"Trainable LoRA params: {sum(p.numel() for p in lora_params):,}")

if GRADIENT_CHECKPOINTING:
    unet.enable_gradient_checkpointing()

# 5. البيانات والـ Optimizer
dataset = TextImageDataset(DATASET_DIR, resolution=RESOLUTION)
print(f"Found {len(dataset)} images in {DATASET_DIR}")
dataloader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True)
optimizer = torch.optim.AdamW(lora_params, lr=LEARNING_RATE)
scaler = torch.amp.GradScaler("cuda", enabled=USE_GRAD_SCALER)

os.makedirs(OUTPUT_DIR, exist_ok=True)


def save_lora(save_dir: str) -> None:
    os.makedirs(save_dir, exist_ok=True)
    state_dict = convert_state_dict_to_diffusers(get_peft_model_state_dict(unet))
    StableDiffusionPipeline.save_lora_weights(
        save_directory=save_dir,
        unet_lora_layers=state_dict,
        weight_name=OUTPUT_NAME,
        safe_serialization=True,
    )


# 6. حلقة التدريب
unet.train()
step = 0
recent_losses = []

for epoch in range(NUM_EPOCHS):
    for batch in dataloader:
        images = batch["pixel_values"].to(DEVICE, dtype=WEIGHT_DTYPE)

        with torch.no_grad():
            latents = vae.encode(images).latent_dist.sample() * vae.config.scaling_factor
            input_ids = tokenizer(
                batch["caption"],
                padding="max_length",
                max_length=tokenizer.model_max_length,
                truncation=True,
                return_tensors="pt",
            ).input_ids.to(DEVICE)
            encoder_hidden_states = text_encoder(input_ids)[0]

        noise = torch.randn_like(latents)
        timesteps = torch.randint(
            0,
            noise_scheduler.config.num_train_timesteps,
            (latents.shape[0],),
            device=DEVICE,
        ).long()
        noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps)

        with torch.autocast("cuda", dtype=AMP_DTYPE, enabled=USE_AMP):
            noise_pred = unet(noisy_latents, timesteps, encoder_hidden_states).sample
        loss = F.mse_loss(noise_pred.float(), noise.float())

        optimizer.zero_grad(set_to_none=True)
        scaler.scale(loss).backward()
        scaler.step(optimizer)
        scaler.update()

        step += 1
        recent_losses.append(loss.item())
        if step % 10 == 0:
            avg = sum(recent_losses) / len(recent_losses)
            print(f"epoch {epoch + 1}/{NUM_EPOCHS}  step {step}  loss {avg:.4f}")
            recent_losses.clear()

        if step % SAVE_EVERY == 0:
            ckpt_dir = os.path.join(OUTPUT_DIR, f"checkpoint-{step}")
            save_lora(ckpt_dir)
            print(f"Saved checkpoint at step {step} -> {ckpt_dir}")

# 7. الحفظ النهائي
print("\nSaving final LoRA weights...")
save_lora(OUTPUT_DIR)
print(f"LoRA saved to: {OUTPUT_DIR}/{OUTPUT_NAME}")



pipe = StableDiffusionPipeline.from_pretrained(
    "stable-diffusion-v1-5/stable-diffusion-v1-5",
    torch_dtype=torch.float16,
    safety_checker=None,
    requires_safety_checker=False,
).to("cuda")


pipe.load_lora_weights("output_lora_test", weight_name="postapo_icon_lora_test.safetensors")

image = pipe("postapo_style, 2d, a rusty gas mask, isolated on white", num_inference_steps=40).images[0]
image.save("test.png")

But I got this image

Image

thanks and regards

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions