updates
This commit is contained in:
@@ -1,28 +0,0 @@
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import os
|
||||
|
||||
def generate_random_color_image(width, height):
|
||||
"""Generate an image of random color."""
|
||||
color = np.random.randint(0, 256, (3,), dtype=np.uint8)
|
||||
image = np.full((height, width, 3), color, dtype=np.uint8)
|
||||
return Image.fromarray(image)
|
||||
|
||||
def save_images(num_images, width, height, directory):
|
||||
"""Save a specified number of random color images."""
|
||||
for i in range(num_images):
|
||||
image = generate_random_color_image(width, height)
|
||||
image.save(f"{directory}/random_color_image_{i+1}.png")
|
||||
|
||||
# Parameters
|
||||
num_images = 40
|
||||
width = 1024
|
||||
height = 1024
|
||||
directory = "random_images"
|
||||
|
||||
os.makedirs(directory, exist_ok=True)
|
||||
|
||||
# Generate and save images
|
||||
save_images(num_images, width, height, directory)
|
||||
|
||||
print(f'Saved {num_images} random color images to {os.path.abspath(directory)}')
|
||||
@@ -291,6 +291,11 @@ def train(config: TrainingConfig):
|
||||
mask = mask.to(config.device)
|
||||
|
||||
captions = list(captions)
|
||||
if config.caption_dropout > 0.0:
|
||||
for i in range(len(captions)):
|
||||
if np.random.rand() < config.caption_dropout:
|
||||
captions[i] = ""
|
||||
|
||||
prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals(
|
||||
config, pipe, captions
|
||||
)
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
"sd_model_version": "sdxl",
|
||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip",
|
||||
"concept_mode": "face",
|
||||
"sample_imgs_lora_scale": 0.75,
|
||||
"sample_imgs_lora_scale": 0.8,
|
||||
"seed": 0,
|
||||
"resolution": 512,
|
||||
"train_batch_size": 4,
|
||||
@@ -11,8 +11,10 @@
|
||||
"max_train_steps": 250,
|
||||
"token_warmup_steps": 0,
|
||||
"checkpointing_steps": 175,
|
||||
|
||||
"disable_ti": false,
|
||||
|
||||
"ti_lr": 0.001,
|
||||
"ti_lr": 0.0005,
|
||||
"unet_lr": 0.001,
|
||||
"lora_rank": 16,
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"name": "twisting_realities",
|
||||
"name": "clipx",
|
||||
"sd_model_version": "sdxl",
|
||||
"lora_training_urls": "https://edenartlab-lfs.s3.amazonaws.com/datasets/twisting_realities.zip",
|
||||
"lora_training_urls": "https://edenartlab-lfs.s3.amazonaws.com/datasets/clipx.zip",
|
||||
"concept_mode": "style",
|
||||
"sample_imgs_lora_scale": 0.8,
|
||||
"seed": 0,
|
||||
|
||||
@@ -38,6 +38,7 @@ class TrainingConfig(BaseModel):
|
||||
concept_mode: Literal["face", "style", "object"]
|
||||
caption_prefix: str = "" # hardcoding this will inject TOK manually and skip the chatgpt token injection step, not recommended unless you know what you're doing
|
||||
caption_model: Literal["gpt4-v", "blip", "florence"] = "blip"
|
||||
caption_dropout: float = 0.0 # dropout rate for captions: occasionally use empty prompt
|
||||
sd_model_version: Literal["sdxl", "sd15", None] = None
|
||||
ckpt_path: str = None # optional hardcoded checkpoint path
|
||||
pretrained_model: dict = None
|
||||
|
||||
+7
-1
@@ -37,7 +37,10 @@ def compute_token_attention_loss(pipe, embedding_handler, captions, masks, daam_
|
||||
att_L2_loss = (torch.relu(mean_att_per_token - att_reg_threshold)**2).mean()
|
||||
att_L2_losses.append(att_L2_loss)
|
||||
|
||||
ti_token_indices = [token_indices.index(token_id) for token_id in embedding_handler.train_ids]
|
||||
try:
|
||||
ti_token_indices = [token_indices.index(token_id) for token_id in embedding_handler.train_ids]
|
||||
except:
|
||||
continue
|
||||
batch_ti_heatmaps, batch_ti_masks = [], []
|
||||
# Extract the attention heatmaps corresponding to the trainable token embeddings:
|
||||
for text_token_index in ti_token_indices:
|
||||
@@ -49,6 +52,9 @@ def compute_token_attention_loss(pipe, embedding_handler, captions, masks, daam_
|
||||
ti_heatmaps.append(torch.stack(batch_ti_heatmaps))
|
||||
ti_masks.append(torch.stack(batch_ti_masks))
|
||||
|
||||
if len(ti_heatmaps) == 0:
|
||||
return torch.tensor(0.0).to(masks.dtype)
|
||||
|
||||
ti_heatmaps = torch.stack(ti_heatmaps)
|
||||
ti_masks = torch.stack(ti_masks)
|
||||
|
||||
|
||||
@@ -17,67 +17,70 @@ import numpy as np
|
||||
from mpl_toolkits.axes_grid1 import make_axes_locatable
|
||||
|
||||
def plot_token_attention_loss(folder, pipe, daam_loss, captions, timesteps, token_attention_loss, global_step, img_ratio):
|
||||
batch_index = 0
|
||||
timestep = timesteps[batch_index].item()
|
||||
if timestep > 700:
|
||||
batch_index += 1
|
||||
try:
|
||||
batch_index = 0
|
||||
timestep = timesteps[batch_index].item()
|
||||
if timestep > 700:
|
||||
batch_index += 1
|
||||
timestep = timesteps[batch_index].item()
|
||||
|
||||
folder = os.path.join(folder, "attention_heatmaps")
|
||||
os.makedirs(folder, exist_ok=True)
|
||||
folder = os.path.join(folder, "attention_heatmaps")
|
||||
os.makedirs(folder, exist_ok=True)
|
||||
|
||||
token_strings = [pipe.tokenizer.decode(x) for x in pipe.tokenizer.encode(captions[batch_index])]
|
||||
plot_token_indices = range(1, len(token_strings) - 1) # Exclude start and end tokens
|
||||
token_strings = [pipe.tokenizer.decode(x) for x in pipe.tokenizer.encode(captions[batch_index])]
|
||||
plot_token_indices = range(1, len(token_strings) - 1) # Exclude start and end tokens
|
||||
|
||||
# Calculate global min and max for consistent colormap
|
||||
all_heatmaps = [daam_loss.get_the_daam_heatmap(text_token_index=i, img_ratio=img_ratio)[batch_index].cpu().detach().float()
|
||||
for i in plot_token_indices]
|
||||
vmin, vmax = np.min([h.min() for h in all_heatmaps]), np.max([h.max() for h in all_heatmaps])
|
||||
# Calculate global min and max for consistent colormap
|
||||
all_heatmaps = [daam_loss.get_the_daam_heatmap(text_token_index=i, img_ratio=img_ratio)[batch_index].cpu().detach().float()
|
||||
for i in plot_token_indices]
|
||||
vmin, vmax = np.min([h.min() for h in all_heatmaps]), np.max([h.max() for h in all_heatmaps])
|
||||
|
||||
# Heatmap plots
|
||||
fig, axes = plt.subplots(nrows=1, ncols=len(plot_token_indices),
|
||||
figsize=(3 * len(plot_token_indices), 10))
|
||||
title_str = f"Token Attention Heatmaps (Step: {global_step})\nDenoise Timestep: {timesteps[batch_index].item()}"
|
||||
fig.suptitle(title_str, fontsize=16)
|
||||
|
||||
for idx, text_token_index in enumerate(plot_token_indices):
|
||||
heatmap = all_heatmaps[idx]
|
||||
im = axes[idx].imshow(heatmap, cmap='viridis', vmin=vmin, vmax=vmax)
|
||||
axes[idx].set_title(f"{token_strings[text_token_index]}")
|
||||
axes[idx].axis("off")
|
||||
# Heatmap plots
|
||||
fig, axes = plt.subplots(nrows=1, ncols=len(plot_token_indices),
|
||||
figsize=(3 * len(plot_token_indices), 10))
|
||||
title_str = f"Token Attention Heatmaps (Step: {global_step})\nDenoise Timestep: {timesteps[batch_index].item()}"
|
||||
fig.suptitle(title_str, fontsize=16)
|
||||
|
||||
# Add colorbar
|
||||
divider = make_axes_locatable(axes[idx])
|
||||
cax = divider.append_axes("bottom", size="5%", pad=0.05)
|
||||
plt.colorbar(im, cax=cax, orientation="horizontal")
|
||||
for idx, text_token_index in enumerate(plot_token_indices):
|
||||
heatmap = all_heatmaps[idx]
|
||||
im = axes[idx].imshow(heatmap, cmap='viridis', vmin=vmin, vmax=vmax)
|
||||
axes[idx].set_title(f"{token_strings[text_token_index]}")
|
||||
axes[idx].axis("off")
|
||||
|
||||
# Add colorbar
|
||||
divider = make_axes_locatable(axes[idx])
|
||||
cax = divider.append_axes("bottom", size="5%", pad=0.05)
|
||||
plt.colorbar(im, cax=cax, orientation="horizontal")
|
||||
|
||||
plt.tight_layout()
|
||||
fig.savefig(os.path.join(folder, f"heatmaps_{global_step}.jpg"), dpi=300, bbox_inches='tight')
|
||||
plt.close(fig)
|
||||
plt.tight_layout()
|
||||
fig.savefig(os.path.join(folder, f"heatmaps_{global_step}.jpg"), dpi=300, bbox_inches='tight')
|
||||
plt.close(fig)
|
||||
|
||||
# Histogram plot
|
||||
fig, ax = plt.subplots(figsize=(12, 6))
|
||||
ax.set_title(f"Token Attention Distribution (Step: {global_step})", fontsize=16)
|
||||
ax.set_xlabel("Attention Value", fontsize=12)
|
||||
ax.set_ylabel("Frequency", fontsize=12)
|
||||
# Histogram plot
|
||||
fig, ax = plt.subplots(figsize=(12, 6))
|
||||
ax.set_title(f"Token Attention Distribution (Step: {global_step})", fontsize=16)
|
||||
ax.set_xlabel("Attention Value", fontsize=12)
|
||||
ax.set_ylabel("Frequency", fontsize=12)
|
||||
|
||||
for idx, text_token_index in enumerate(plot_token_indices):
|
||||
heatmap = all_heatmaps[idx]
|
||||
ax.hist(heatmap.reshape(-1), bins=30, label=token_strings[text_token_index],
|
||||
alpha=0.5, density=True)
|
||||
for idx, text_token_index in enumerate(plot_token_indices):
|
||||
heatmap = all_heatmaps[idx]
|
||||
ax.hist(heatmap.reshape(-1), bins=30, label=token_strings[text_token_index],
|
||||
alpha=0.5, density=True)
|
||||
|
||||
ax.legend(bbox_to_anchor=(1.05, 1), loc='upper left', fontsize=10)
|
||||
ax.grid(alpha=0.3)
|
||||
plt.tight_layout()
|
||||
ax.legend(bbox_to_anchor=(1.05, 1), loc='upper left', fontsize=10)
|
||||
ax.grid(alpha=0.3)
|
||||
plt.tight_layout()
|
||||
|
||||
# Add token attention loss as text
|
||||
plt.text(0.95, 0.95, f"Token Attention Loss: {token_attention_loss.item():.4f}",
|
||||
transform=ax.transAxes, ha='right', va='top',
|
||||
bbox=dict(facecolor='white', edgecolor='black', alpha=0.8))
|
||||
# Add token attention loss as text
|
||||
plt.text(0.95, 0.95, f"Token Attention Loss: {token_attention_loss.item():.4f}",
|
||||
transform=ax.transAxes, ha='right', va='top',
|
||||
bbox=dict(facecolor='white', edgecolor='black', alpha=0.8))
|
||||
|
||||
plt.ylim(0, 0.6)
|
||||
fig.savefig(os.path.join(folder, f"histogram_{global_step}.jpg"), dpi=300, bbox_inches='tight')
|
||||
plt.close(fig)
|
||||
plt.ylim(0, 0.6)
|
||||
fig.savefig(os.path.join(folder, f"histogram_{global_step}.jpg"), dpi=300, bbox_inches='tight')
|
||||
plt.close(fig)
|
||||
except:
|
||||
print("Failed to plot token attention loss")
|
||||
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user