comdoleger commited on
Commit
621e3a2
·
verified ·
1 Parent(s): dd53e76

Upload extensions_built_in/ultimate_slider_trainer/UltimateSliderTrainerProcess.py with huggingface_hub

Browse files
extensions_built_in/ultimate_slider_trainer/UltimateSliderTrainerProcess.py ADDED
@@ -0,0 +1,533 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import copy
2
+ import random
3
+ from collections import OrderedDict
4
+ import os
5
+ from contextlib import nullcontext
6
+ from typing import Optional, Union, List
7
+ from torch.utils.data import ConcatDataset, DataLoader
8
+
9
+ from toolkit.config_modules import ReferenceDatasetConfig
10
+ from toolkit.data_loader import PairedImageDataset
11
+ from toolkit.prompt_utils import concat_prompt_embeds, split_prompt_embeds, build_latent_image_batch_for_prompt_pair
12
+ from toolkit.stable_diffusion_model import StableDiffusion, PromptEmbeds
13
+ from toolkit.train_tools import get_torch_dtype, apply_snr_weight
14
+ import gc
15
+ from toolkit import train_tools
16
+ import torch
17
+ from jobs.process import BaseSDTrainProcess
18
+ import random
19
+
20
+ import random
21
+ from collections import OrderedDict
22
+ from tqdm import tqdm
23
+
24
+ from toolkit.config_modules import SliderConfig
25
+ from toolkit.train_tools import get_torch_dtype, apply_snr_weight
26
+ import gc
27
+ from toolkit import train_tools
28
+ from toolkit.prompt_utils import \
29
+ EncodedPromptPair, ACTION_TYPES_SLIDER, \
30
+ EncodedAnchor, concat_prompt_pairs, \
31
+ concat_anchors, PromptEmbedsCache, encode_prompts_to_cache, build_prompt_pair_batch_from_cache, split_anchors, \
32
+ split_prompt_pairs
33
+
34
+ import torch
35
+
36
+
37
+ def flush():
38
+ torch.cuda.empty_cache()
39
+ gc.collect()
40
+
41
+
42
+ class UltimateSliderConfig(SliderConfig):
43
+ def __init__(self, **kwargs):
44
+ super().__init__(**kwargs)
45
+ self.additional_losses: List[str] = kwargs.get('additional_losses', [])
46
+ self.weight_jitter: float = kwargs.get('weight_jitter', 0.0)
47
+ self.img_loss_weight: float = kwargs.get('img_loss_weight', 1.0)
48
+ self.cfg_loss_weight: float = kwargs.get('cfg_loss_weight', 1.0)
49
+ self.datasets: List[ReferenceDatasetConfig] = [ReferenceDatasetConfig(**d) for d in kwargs.get('datasets', [])]
50
+
51
+
52
+ class UltimateSliderTrainerProcess(BaseSDTrainProcess):
53
+ sd: StableDiffusion
54
+ data_loader: DataLoader = None
55
+
56
+ def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
57
+ super().__init__(process_id, job, config, **kwargs)
58
+ self.prompt_txt_list = None
59
+ self.step_num = 0
60
+ self.start_step = 0
61
+ self.device = self.get_conf('device', self.job.device)
62
+ self.device_torch = torch.device(self.device)
63
+ self.slider_config = UltimateSliderConfig(**self.get_conf('slider', {}))
64
+
65
+ self.prompt_cache = PromptEmbedsCache()
66
+ self.prompt_pairs: list[EncodedPromptPair] = []
67
+ self.anchor_pairs: list[EncodedAnchor] = []
68
+ # keep track of prompt chunk size
69
+ self.prompt_chunk_size = 1
70
+
71
+ # store a list of all the prompts from the dataset so we can cache it
72
+ self.dataset_prompts = []
73
+ self.train_with_dataset = self.slider_config.datasets is not None and len(self.slider_config.datasets) > 0
74
+
75
+ def load_datasets(self):
76
+ if self.data_loader is None and \
77
+ self.slider_config.datasets is not None and len(self.slider_config.datasets) > 0:
78
+ print(f"Loading datasets")
79
+ datasets = []
80
+ for dataset in self.slider_config.datasets:
81
+ print(f" - Dataset: {dataset.pair_folder}")
82
+ config = {
83
+ 'path': dataset.pair_folder,
84
+ 'size': dataset.size,
85
+ 'default_prompt': dataset.target_class,
86
+ 'network_weight': dataset.network_weight,
87
+ 'pos_weight': dataset.pos_weight,
88
+ 'neg_weight': dataset.neg_weight,
89
+ 'pos_folder': dataset.pos_folder,
90
+ 'neg_folder': dataset.neg_folder,
91
+ }
92
+ image_dataset = PairedImageDataset(config)
93
+ datasets.append(image_dataset)
94
+
95
+ # capture all the prompts from it so we can cache the embeds
96
+ self.dataset_prompts += image_dataset.get_all_prompts()
97
+
98
+ concatenated_dataset = ConcatDataset(datasets)
99
+ self.data_loader = DataLoader(
100
+ concatenated_dataset,
101
+ batch_size=self.train_config.batch_size,
102
+ shuffle=True,
103
+ num_workers=2
104
+ )
105
+
106
+ def before_model_load(self):
107
+ pass
108
+
109
+ def hook_before_train_loop(self):
110
+ # load any datasets if they were passed
111
+ self.load_datasets()
112
+
113
+ # read line by line from file
114
+ if self.slider_config.prompt_file:
115
+ self.print(f"Loading prompt file from {self.slider_config.prompt_file}")
116
+ with open(self.slider_config.prompt_file, 'r', encoding='utf-8') as f:
117
+ self.prompt_txt_list = f.readlines()
118
+ # clean empty lines
119
+ self.prompt_txt_list = [line.strip() for line in self.prompt_txt_list if len(line.strip()) > 0]
120
+
121
+ self.print(f"Found {len(self.prompt_txt_list)} prompts.")
122
+
123
+ if not self.slider_config.prompt_tensors:
124
+ print(f"Prompt tensors not found. Building prompt tensors for {self.train_config.steps} steps.")
125
+ # shuffle
126
+ random.shuffle(self.prompt_txt_list)
127
+ # trim to max steps
128
+ self.prompt_txt_list = self.prompt_txt_list[:self.train_config.steps]
129
+ # trim list to our max steps
130
+
131
+ cache = PromptEmbedsCache()
132
+
133
+ # get encoded latents for our prompts
134
+ with torch.no_grad():
135
+ # list of neutrals. Can come from file or be empty
136
+ neutral_list = self.prompt_txt_list if self.prompt_txt_list is not None else [""]
137
+
138
+ # build the prompts to cache
139
+ prompts_to_cache = []
140
+ for neutral in neutral_list:
141
+ for target in self.slider_config.targets:
142
+ prompt_list = [
143
+ f"{target.target_class}", # target_class
144
+ f"{target.target_class} {neutral}", # target_class with neutral
145
+ f"{target.positive}", # positive_target
146
+ f"{target.positive} {neutral}", # positive_target with neutral
147
+ f"{target.negative}", # negative_target
148
+ f"{target.negative} {neutral}", # negative_target with neutral
149
+ f"{neutral}", # neutral
150
+ f"{target.positive} {target.negative}", # both targets
151
+ f"{target.negative} {target.positive}", # both targets reverse
152
+ ]
153
+ prompts_to_cache += prompt_list
154
+
155
+ # remove duplicates
156
+ prompts_to_cache = list(dict.fromkeys(prompts_to_cache))
157
+
158
+ # trim to max steps if max steps is lower than prompt count
159
+ prompts_to_cache = prompts_to_cache[:self.train_config.steps]
160
+
161
+ if len(self.dataset_prompts) > 0:
162
+ # add the prompts from the dataset
163
+ prompts_to_cache += self.dataset_prompts
164
+
165
+ # encode them
166
+ cache = encode_prompts_to_cache(
167
+ prompt_list=prompts_to_cache,
168
+ sd=self.sd,
169
+ cache=cache,
170
+ prompt_tensor_file=self.slider_config.prompt_tensors
171
+ )
172
+
173
+ prompt_pairs = []
174
+ prompt_batches = []
175
+ for neutral in tqdm(neutral_list, desc="Building Prompt Pairs", leave=False):
176
+ for target in self.slider_config.targets:
177
+ prompt_pair_batch = build_prompt_pair_batch_from_cache(
178
+ cache=cache,
179
+ target=target,
180
+ neutral=neutral,
181
+
182
+ )
183
+ if self.slider_config.batch_full_slide:
184
+ # concat the prompt pairs
185
+ # this allows us to run the entire 4 part process in one shot (for slider)
186
+ self.prompt_chunk_size = 4
187
+ concat_prompt_pair_batch = concat_prompt_pairs(prompt_pair_batch).to('cpu')
188
+ prompt_pairs += [concat_prompt_pair_batch]
189
+ else:
190
+ self.prompt_chunk_size = 1
191
+ # do them one at a time (probably not necessary after new optimizations)
192
+ prompt_pairs += [x.to('cpu') for x in prompt_pair_batch]
193
+
194
+ # move to cpu to save vram
195
+ # We don't need text encoder anymore, but keep it on cpu for sampling
196
+ # if text encoder is list
197
+ if isinstance(self.sd.text_encoder, list):
198
+ for encoder in self.sd.text_encoder:
199
+ encoder.to("cpu")
200
+ else:
201
+ self.sd.text_encoder.to("cpu")
202
+ self.prompt_cache = cache
203
+ self.prompt_pairs = prompt_pairs
204
+ # end hook_before_train_loop
205
+
206
+ # move vae to device so we can encode on the fly
207
+ # todo cache latents
208
+ self.sd.vae.to(self.device_torch)
209
+ self.sd.vae.eval()
210
+ self.sd.vae.requires_grad_(False)
211
+
212
+ if self.train_config.gradient_checkpointing:
213
+ # may get disabled elsewhere
214
+ self.sd.unet.enable_gradient_checkpointing()
215
+
216
+ flush()
217
+ # end hook_before_train_loop
218
+
219
+ def hook_train_loop(self, batch):
220
+ dtype = get_torch_dtype(self.train_config.dtype)
221
+
222
+ with torch.no_grad():
223
+ ### LOOP SETUP ###
224
+ noise_scheduler = self.sd.noise_scheduler
225
+ optimizer = self.optimizer
226
+ lr_scheduler = self.lr_scheduler
227
+
228
+ ### TARGET_PROMPTS ###
229
+ # get a random pair
230
+ prompt_pair: EncodedPromptPair = self.prompt_pairs[
231
+ torch.randint(0, len(self.prompt_pairs), (1,)).item()
232
+ ]
233
+ # move to device and dtype
234
+ prompt_pair.to(self.device_torch, dtype=dtype)
235
+
236
+ ### PREP REFERENCE IMAGES ###
237
+
238
+ imgs, prompts, network_weights = batch
239
+ network_pos_weight, network_neg_weight = network_weights
240
+
241
+ if isinstance(network_pos_weight, torch.Tensor):
242
+ network_pos_weight = network_pos_weight.item()
243
+ if isinstance(network_neg_weight, torch.Tensor):
244
+ network_neg_weight = network_neg_weight.item()
245
+
246
+ # get an array of random floats between -weight_jitter and weight_jitter
247
+ weight_jitter = self.slider_config.weight_jitter
248
+ if weight_jitter > 0.0:
249
+ jitter_list = random.uniform(-weight_jitter, weight_jitter)
250
+ network_pos_weight += jitter_list
251
+ network_neg_weight += (jitter_list * -1.0)
252
+
253
+ # if items in network_weight list are tensors, convert them to floats
254
+ imgs: torch.Tensor = imgs.to(self.device_torch, dtype=dtype)
255
+ # split batched images in half so left is negative and right is positive
256
+ negative_images, positive_images = torch.chunk(imgs, 2, dim=3)
257
+
258
+ height = positive_images.shape[2]
259
+ width = positive_images.shape[3]
260
+ batch_size = positive_images.shape[0]
261
+
262
+ positive_latents = self.sd.encode_images(positive_images)
263
+ negative_latents = self.sd.encode_images(negative_images)
264
+
265
+ self.sd.noise_scheduler.set_timesteps(
266
+ self.train_config.max_denoising_steps, device=self.device_torch
267
+ )
268
+
269
+ timesteps = torch.randint(0, self.train_config.max_denoising_steps, (1,), device=self.device_torch)
270
+ current_timestep_index = timesteps.item()
271
+ current_timestep = noise_scheduler.timesteps[current_timestep_index]
272
+ timesteps = timesteps.long()
273
+
274
+ # get noise
275
+ noise_positive = self.sd.get_latent_noise(
276
+ pixel_height=height,
277
+ pixel_width=width,
278
+ batch_size=batch_size,
279
+ noise_offset=self.train_config.noise_offset,
280
+ ).to(self.device_torch, dtype=dtype)
281
+
282
+ noise_negative = noise_positive.clone()
283
+
284
+ # Add noise to the latents according to the noise magnitude at each timestep
285
+ # (this is the forward diffusion process)
286
+ noisy_positive_latents = noise_scheduler.add_noise(positive_latents, noise_positive, timesteps)
287
+ noisy_negative_latents = noise_scheduler.add_noise(negative_latents, noise_negative, timesteps)
288
+
289
+ ### CFG SLIDER TRAINING PREP ###
290
+
291
+ # get CFG txt latents
292
+ noisy_cfg_latents = build_latent_image_batch_for_prompt_pair(
293
+ pos_latent=noisy_positive_latents,
294
+ neg_latent=noisy_negative_latents,
295
+ prompt_pair=prompt_pair,
296
+ prompt_chunk_size=self.prompt_chunk_size,
297
+ )
298
+ noisy_cfg_latents.requires_grad = False
299
+
300
+ assert not self.network.is_active
301
+
302
+ # 4.20 GB RAM for 512x512
303
+ positive_latents = self.sd.predict_noise(
304
+ latents=noisy_cfg_latents,
305
+ text_embeddings=train_tools.concat_prompt_embeddings(
306
+ prompt_pair.positive_target, # negative prompt
307
+ prompt_pair.negative_target, # positive prompt
308
+ self.train_config.batch_size,
309
+ ),
310
+ timestep=current_timestep,
311
+ guidance_scale=1.0
312
+ )
313
+ positive_latents.requires_grad = False
314
+
315
+ neutral_latents = self.sd.predict_noise(
316
+ latents=noisy_cfg_latents,
317
+ text_embeddings=train_tools.concat_prompt_embeddings(
318
+ prompt_pair.positive_target, # negative prompt
319
+ prompt_pair.empty_prompt, # positive prompt (normally neutral
320
+ self.train_config.batch_size,
321
+ ),
322
+ timestep=current_timestep,
323
+ guidance_scale=1.0
324
+ )
325
+ neutral_latents.requires_grad = False
326
+
327
+ unconditional_latents = self.sd.predict_noise(
328
+ latents=noisy_cfg_latents,
329
+ text_embeddings=train_tools.concat_prompt_embeddings(
330
+ prompt_pair.positive_target, # negative prompt
331
+ prompt_pair.positive_target, # positive prompt
332
+ self.train_config.batch_size,
333
+ ),
334
+ timestep=current_timestep,
335
+ guidance_scale=1.0
336
+ )
337
+ unconditional_latents.requires_grad = False
338
+
339
+ positive_latents_chunks = torch.chunk(positive_latents, self.prompt_chunk_size, dim=0)
340
+ neutral_latents_chunks = torch.chunk(neutral_latents, self.prompt_chunk_size, dim=0)
341
+ unconditional_latents_chunks = torch.chunk(unconditional_latents, self.prompt_chunk_size, dim=0)
342
+ prompt_pair_chunks = split_prompt_pairs(prompt_pair, self.prompt_chunk_size)
343
+ noisy_cfg_latents_chunks = torch.chunk(noisy_cfg_latents, self.prompt_chunk_size, dim=0)
344
+ assert len(prompt_pair_chunks) == len(noisy_cfg_latents_chunks)
345
+
346
+ noisy_latents = torch.cat([noisy_positive_latents, noisy_negative_latents], dim=0)
347
+ noise = torch.cat([noise_positive, noise_negative], dim=0)
348
+ timesteps = torch.cat([timesteps, timesteps], dim=0)
349
+ network_multiplier = [network_pos_weight * 1.0, network_neg_weight * -1.0]
350
+
351
+ flush()
352
+
353
+ loss_float = None
354
+ loss_mirror_float = None
355
+
356
+ self.optimizer.zero_grad()
357
+ noisy_latents.requires_grad = False
358
+
359
+ # TODO allow both processed to train text encoder, for now, we just to unet and cache all text encodes
360
+ # if training text encoder enable grads, else do context of no grad
361
+ # with torch.set_grad_enabled(self.train_config.train_text_encoder):
362
+ # # text encoding
363
+ # embedding_list = []
364
+ # # embed the prompts
365
+ # for prompt in prompts:
366
+ # embedding = self.sd.encode_prompt(prompt).to(self.device_torch, dtype=dtype)
367
+ # embedding_list.append(embedding)
368
+ # conditional_embeds = concat_prompt_embeds(embedding_list)
369
+ # conditional_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
370
+
371
+ if self.train_with_dataset:
372
+ embedding_list = []
373
+ with torch.set_grad_enabled(self.train_config.train_text_encoder):
374
+ for prompt in prompts:
375
+ # get embedding form cache
376
+ embedding = self.prompt_cache[prompt]
377
+ embedding = embedding.to(self.device_torch, dtype=dtype)
378
+ embedding_list.append(embedding)
379
+ conditional_embeds = concat_prompt_embeds(embedding_list)
380
+ # double up so we can do both sides of the slider
381
+ conditional_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
382
+ else:
383
+ # throw error. Not supported yet
384
+ raise Exception("Datasets and targets required for ultimate slider")
385
+
386
+ if self.model_config.is_xl:
387
+ # todo also allow for setting this for low ram in general, but sdxl spikes a ton on back prop
388
+ network_multiplier_list = network_multiplier
389
+ noisy_latent_list = torch.chunk(noisy_latents, 2, dim=0)
390
+ noise_list = torch.chunk(noise, 2, dim=0)
391
+ timesteps_list = torch.chunk(timesteps, 2, dim=0)
392
+ conditional_embeds_list = split_prompt_embeds(conditional_embeds)
393
+ else:
394
+ network_multiplier_list = [network_multiplier]
395
+ noisy_latent_list = [noisy_latents]
396
+ noise_list = [noise]
397
+ timesteps_list = [timesteps]
398
+ conditional_embeds_list = [conditional_embeds]
399
+
400
+ ## DO REFERENCE IMAGE TRAINING ##
401
+
402
+ reference_image_losses = []
403
+ # allow to chunk it out to save vram
404
+ for network_multiplier, noisy_latents, noise, timesteps, conditional_embeds in zip(
405
+ network_multiplier_list, noisy_latent_list, noise_list, timesteps_list, conditional_embeds_list
406
+ ):
407
+ with self.network:
408
+ assert self.network.is_active
409
+
410
+ self.network.multiplier = network_multiplier
411
+
412
+ noise_pred = self.sd.predict_noise(
413
+ latents=noisy_latents.to(self.device_torch, dtype=dtype),
414
+ conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
415
+ timestep=timesteps,
416
+ )
417
+ noise = noise.to(self.device_torch, dtype=dtype)
418
+
419
+ if self.sd.prediction_type == 'v_prediction':
420
+ # v-parameterization training
421
+ target = noise_scheduler.get_velocity(noisy_latents, noise, timesteps)
422
+ else:
423
+ target = noise
424
+
425
+ loss = torch.nn.functional.mse_loss(noise_pred.float(), target.float(), reduction="none")
426
+ loss = loss.mean([1, 2, 3])
427
+
428
+ # todo add snr gamma here
429
+ if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
430
+ # add min_snr_gamma
431
+ loss = apply_snr_weight(loss, timesteps, noise_scheduler, self.train_config.min_snr_gamma)
432
+
433
+ loss = loss.mean()
434
+ loss = loss * self.slider_config.img_loss_weight
435
+ loss_slide_float = loss.item()
436
+
437
+ loss_float = loss.item()
438
+ reference_image_losses.append(loss_float)
439
+
440
+ # back propagate loss to free ram
441
+ loss.backward()
442
+ flush()
443
+
444
+ ## DO CFG SLIDER TRAINING ##
445
+
446
+ cfg_loss_list = []
447
+
448
+ with self.network:
449
+ assert self.network.is_active
450
+ for prompt_pair_chunk, \
451
+ noisy_cfg_latent_chunk, \
452
+ positive_latents_chunk, \
453
+ neutral_latents_chunk, \
454
+ unconditional_latents_chunk \
455
+ in zip(
456
+ prompt_pair_chunks,
457
+ noisy_cfg_latents_chunks,
458
+ positive_latents_chunks,
459
+ neutral_latents_chunks,
460
+ unconditional_latents_chunks,
461
+ ):
462
+ self.network.multiplier = prompt_pair_chunk.multiplier_list
463
+
464
+ target_latents = self.sd.predict_noise(
465
+ latents=noisy_cfg_latent_chunk,
466
+ text_embeddings=train_tools.concat_prompt_embeddings(
467
+ prompt_pair_chunk.positive_target, # negative prompt
468
+ prompt_pair_chunk.target_class, # positive prompt
469
+ self.train_config.batch_size,
470
+ ),
471
+ timestep=current_timestep,
472
+ guidance_scale=1.0
473
+ )
474
+
475
+ guidance_scale = 1.0
476
+
477
+ offset = guidance_scale * (positive_latents_chunk - unconditional_latents_chunk)
478
+
479
+ # make offset multiplier based on actions
480
+ offset_multiplier_list = []
481
+ for action in prompt_pair_chunk.action_list:
482
+ if action == ACTION_TYPES_SLIDER.ERASE_NEGATIVE:
483
+ offset_multiplier_list += [-1.0]
484
+ elif action == ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE:
485
+ offset_multiplier_list += [1.0]
486
+
487
+ offset_multiplier = torch.tensor(offset_multiplier_list).to(offset.device, dtype=offset.dtype)
488
+ # make offset multiplier match rank of offset
489
+ offset_multiplier = offset_multiplier.view(offset.shape[0], 1, 1, 1)
490
+ offset *= offset_multiplier
491
+
492
+ offset_neutral = neutral_latents_chunk
493
+ # offsets are already adjusted on a per-batch basis
494
+ offset_neutral += offset
495
+
496
+ # 16.15 GB RAM for 512x512 -> 4.20GB RAM for 512x512 with new grad_checkpointing
497
+ loss = torch.nn.functional.mse_loss(target_latents.float(), offset_neutral.float(), reduction="none")
498
+ loss = loss.mean([1, 2, 3])
499
+
500
+ if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
501
+ # match batch size
502
+ timesteps_index_list = [current_timestep_index for _ in range(target_latents.shape[0])]
503
+ # add min_snr_gamma
504
+ loss = apply_snr_weight(loss, timesteps_index_list, noise_scheduler,
505
+ self.train_config.min_snr_gamma)
506
+
507
+ loss = loss.mean() * prompt_pair_chunk.weight * self.slider_config.cfg_loss_weight
508
+
509
+ loss.backward()
510
+ cfg_loss_list.append(loss.item())
511
+ del target_latents
512
+ del offset_neutral
513
+ del loss
514
+ flush()
515
+
516
+ # apply gradients
517
+ optimizer.step()
518
+ lr_scheduler.step()
519
+
520
+ # reset network
521
+ self.network.multiplier = 1.0
522
+
523
+ reference_image_loss = sum(reference_image_losses) / len(reference_image_losses) if len(
524
+ reference_image_losses) > 0 else 0.0
525
+ cfg_loss = sum(cfg_loss_list) / len(cfg_loss_list) if len(cfg_loss_list) > 0 else 0.0
526
+
527
+ loss_dict = OrderedDict({
528
+ 'loss/img': reference_image_loss,
529
+ 'loss/cfg': cfg_loss,
530
+ })
531
+
532
+ return loss_dict
533
+ # end hook_train_loop