Text-to-Image
Diffusers
Safetensors
recoilme commited on
Commit
fee5085
·
1 Parent(s): a1aa270
girl.jpg CHANGED

Git LFS Details

  • SHA256: d6df0658f6b18b224b6c3c504f7c3663ec22c887bba9b5f587db48db66d8071d
  • Pointer size: 130 Bytes
  • Size of remote file: 99.8 kB

Git LFS Details

  • SHA256: 1a893c9d4e115d1d7b63eff4673ad12aaa5f53159654b2fabdfc8da49b6360c2
  • Pointer size: 130 Bytes
  • Size of remote file: 99.7 kB
media/result_grid.jpg CHANGED

Git LFS Details

  • SHA256: 99b48552aabe20d7a468a1bfaa93822863acda59846c92da9d0ecd85d21fc65f
  • Pointer size: 132 Bytes
  • Size of remote file: 2.98 MB

Git LFS Details

  • SHA256: 6d44813a72a2f716f2158007e58bf7d86f3d70bf6b7a9cec4f130073f14277d7
  • Pointer size: 132 Bytes
  • Size of remote file: 2.99 MB
pipeline_sdxs.py CHANGED
@@ -80,7 +80,7 @@ class SdxsPipeline(DiffusionPipeline):
80
  toks = self.tokenizer(
81
  formatted_prompts,
82
  padding="max_length",
83
- max_length=255,
84
  truncation=True, # Не забываем обрезать, если вдруг длиннее
85
  return_tensors="pt"
86
  ).to(device)
@@ -97,25 +97,34 @@ class SdxsPipeline(DiffusionPipeline):
97
  seq_len = toks.attention_mask.sum(dim=1) - 1
98
  pooled = last_hidden[torch.arange(len(last_hidden)), seq_len.clamp(min=0)]
99
 
100
- return last_hidden, toks.attention_mask, pooled
 
 
 
 
 
 
 
 
 
 
 
 
101
 
102
- pos_embeds, pos_mask, pos_pooled = get_encode(prompt)
103
- neg_embeds, neg_mask, neg_pooled = get_encode(negative_prompt)
 
 
104
 
105
  batch_size = pos_embeds.shape[0]
106
  if neg_embeds.shape[0] != batch_size:
107
  neg_embeds = neg_embeds.repeat(batch_size, 1, 1)
108
  neg_mask = neg_mask.repeat(batch_size, 1)
109
- neg_pooled = neg_pooled.repeat(batch_size, 1)
110
-
111
- if pos_pooled.shape[0] != batch_size:
112
- pos_pooled = pos_pooled.repeat(batch_size, 1)
113
 
114
  text_embeddings = torch.cat([neg_embeds, pos_embeds], dim=0)
115
  final_mask = torch.cat([neg_mask, pos_mask], dim=0)
116
- pooled_embeds = torch.cat([neg_pooled, pos_pooled], dim=0)
117
 
118
- return text_embeddings.to(dtype=dtype), final_mask.to(dtype=torch.int64), pooled_embeds.to(dtype=dtype)
119
 
120
  @torch.no_grad()
121
  def __call__(
@@ -165,7 +174,7 @@ class SdxsPipeline(DiffusionPipeline):
165
  ).to(device)
166
 
167
  generated_ids = self.text_encoder.generate(
168
- **inputs, max_new_tokens=255, do_sample=True,temperature = 0.7
169
  )
170
 
171
  # Обрезаем входные токены из ответа
@@ -180,7 +189,7 @@ class SdxsPipeline(DiffusionPipeline):
180
  prompt = refined_list[0] if isinstance(prompt, str) else refined_list
181
 
182
  # ==================== ENCODE PROMPTS ====================
183
- text_embeddings, attention_mask, pooled_embeds = self.encode_prompt(
184
  prompt, negative_prompt, device, dtype
185
  )
186
  batch_size = 1 if isinstance(prompt, str) else len(prompt)
@@ -188,14 +197,6 @@ class SdxsPipeline(DiffusionPipeline):
188
  # 2. Scheduler timesteps
189
  self.scheduler.set_timesteps(num_inference_steps, device=device)
190
  timesteps = self.scheduler.timesteps
191
-
192
- # ==================== TIME IDS =======================================
193
- time_ids = torch.zeros(
194
- pooled_embeds.shape[0],
195
- 6,
196
- device=device,
197
- dtype=torch.long
198
- )
199
 
200
  # ==================== IMG2IMG БЛОК (НОВАЯ ВЕРСИЯ) ====================
201
  if image is not None:
@@ -244,7 +245,6 @@ class SdxsPipeline(DiffusionPipeline):
244
  t,
245
  encoder_hidden_states=text_embeddings,
246
  encoder_attention_mask=attention_mask,
247
- added_cond_kwargs={"text_embeds": pooled_embeds,"time_ids": time_ids},
248
  return_dict=False,
249
  )[0]
250
 
 
80
  toks = self.tokenizer(
81
  formatted_prompts,
82
  padding="max_length",
83
+ max_length=248,
84
  truncation=True, # Не забываем обрезать, если вдруг длиннее
85
  return_tensors="pt"
86
  ).to(device)
 
97
  seq_len = toks.attention_mask.sum(dim=1) - 1
98
  pooled = last_hidden[torch.arange(len(last_hidden)), seq_len.clamp(min=0)]
99
 
100
+ # --- НОВАЯ ЛОГИКА: ОБЪЕДИНЕНИЕ ДЛЯ КРОСС-ВНИМАНИЯ ---
101
+ # 1. Расширяем пулинг-вектор до последовательности [B, 1, 1024]
102
+ pooled_expanded = pooled.unsqueeze(1)
103
+
104
+ # 2. Объединяем последовательность токенов и пулинг-вектор
105
+ # !!! ИЗМЕНЕНИЕ ЗДЕСЬ !!!: Пулинг идет ПЕРВЫМ
106
+ # Теперь: [B, 1 + L, 1024]. Пулинг стал токеном в НАЧАЛЕ.
107
+ new_encoder_hidden_states = torch.cat([pooled_expanded, last_hidden], dim=1)
108
+
109
+ # 3. Обновляем маску внимания для нового токена
110
+ # Маска внимания: [B, 1 + L]. Добавляем 1 в НАЧАЛО.
111
+ # torch.ones((batch_size, 1), device=device) создает маску [B, 1] со значениями 1.
112
+ new_attention_mask = torch.cat([torch.ones((last_hidden.shape[0], 1), device=device), toks.attention_mask], dim=1)
113
 
114
+ return new_encoder_hidden_states, new_attention_mask
115
+
116
+ pos_embeds, pos_mask = get_encode(prompt)
117
+ neg_embeds, neg_mask = get_encode(negative_prompt)
118
 
119
  batch_size = pos_embeds.shape[0]
120
  if neg_embeds.shape[0] != batch_size:
121
  neg_embeds = neg_embeds.repeat(batch_size, 1, 1)
122
  neg_mask = neg_mask.repeat(batch_size, 1)
 
 
 
 
123
 
124
  text_embeddings = torch.cat([neg_embeds, pos_embeds], dim=0)
125
  final_mask = torch.cat([neg_mask, pos_mask], dim=0)
 
126
 
127
+ return text_embeddings.to(dtype=dtype), final_mask.to(dtype=torch.int64)
128
 
129
  @torch.no_grad()
130
  def __call__(
 
174
  ).to(device)
175
 
176
  generated_ids = self.text_encoder.generate(
177
+ **inputs, max_new_tokens=248, do_sample=True,temperature = 0.7
178
  )
179
 
180
  # Обрезаем входные токены из ответа
 
189
  prompt = refined_list[0] if isinstance(prompt, str) else refined_list
190
 
191
  # ==================== ENCODE PROMPTS ====================
192
+ text_embeddings, attention_mask = self.encode_prompt(
193
  prompt, negative_prompt, device, dtype
194
  )
195
  batch_size = 1 if isinstance(prompt, str) else len(prompt)
 
197
  # 2. Scheduler timesteps
198
  self.scheduler.set_timesteps(num_inference_steps, device=device)
199
  timesteps = self.scheduler.timesteps
 
 
 
 
 
 
 
 
200
 
201
  # ==================== IMG2IMG БЛОК (НОВАЯ ВЕРСИЯ) ====================
202
  if image is not None:
 
245
  t,
246
  encoder_hidden_states=text_embeddings,
247
  encoder_attention_mask=attention_mask,
 
248
  return_dict=False,
249
  )[0]
250
 
samples/unet_1024x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: 5e629a775120e3b41844070de8923bffe09d7d3dfc73811fa6320b43a09f756f
  • Pointer size: 131 Bytes
  • Size of remote file: 385 kB

Git LFS Details

  • SHA256: 15445f3050ed052ff3eef27175a79e3c7b52fa7b7796db4ccca7941f9bf3760b
  • Pointer size: 131 Bytes
  • Size of remote file: 558 kB
samples/unet_1088x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: efb61b977c5f4faf23c2331f138c12f9077fd0eca8fa257c53cd6e0d1ab12a27
  • Pointer size: 131 Bytes
  • Size of remote file: 157 kB

Git LFS Details

  • SHA256: 777c31666a97d073806c2dcca95dc7ab53f517ce7ec7fd6b1ef40009e86218d3
  • Pointer size: 131 Bytes
  • Size of remote file: 225 kB
samples/unet_1152x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: e5b27ad228e4697af4b3f46ff71e04dc80eec88ce5947f41b310fc820210b2eb
  • Pointer size: 131 Bytes
  • Size of remote file: 156 kB

Git LFS Details

  • SHA256: d93520843b1e0c5351d944e8b10f37be9aa2f8a27b0c6a9cc4d20e9865c37563
  • Pointer size: 131 Bytes
  • Size of remote file: 271 kB
samples/unet_1216x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: eaf12941a048f6ab0b045067904e7d5c38b2821ce5293ec2f345d00850053b9a
  • Pointer size: 131 Bytes
  • Size of remote file: 326 kB

Git LFS Details

  • SHA256: b4764c1101d05bcd5afacd84288d65553f46fcca9562499ddc7459c1f3e1e378
  • Pointer size: 131 Bytes
  • Size of remote file: 394 kB
samples/unet_1280x1024_0.jpg CHANGED

Git LFS Details

  • SHA256: 39fb8851ebc9ce1e2d1e089f2b61cce1f8bc1572393b3ec1c9aaf34f256b8e2e
  • Pointer size: 131 Bytes
  • Size of remote file: 215 kB

Git LFS Details

  • SHA256: 2b9e68413411aff989e82b23caf9beeb93689946d2d8ec6e5968b713a50b3bd3
  • Pointer size: 131 Bytes
  • Size of remote file: 107 kB
samples/unet_1280x1088_0.jpg CHANGED

Git LFS Details

  • SHA256: db87f92fcda29d504098544bf6dd77e16e76f006423bdc7e39be316e36407703
  • Pointer size: 131 Bytes
  • Size of remote file: 429 kB

Git LFS Details

  • SHA256: 220863d9f4ad1e3d70ef0f5d4d907428a428f8ab0d2c8055bd8d3252e3e9168b
  • Pointer size: 131 Bytes
  • Size of remote file: 701 kB
samples/unet_1280x1152_0.jpg CHANGED

Git LFS Details

  • SHA256: 4079943d66b918b1a576e597d2c3158d388870fbf702c6c99a317117ecd2b729
  • Pointer size: 131 Bytes
  • Size of remote file: 211 kB

Git LFS Details

  • SHA256: 7cb111f695247f2ec1d4d53b0bb5568289eb979ba33864fc7ebd8ff41b9c4539
  • Pointer size: 131 Bytes
  • Size of remote file: 408 kB
samples/unet_1280x1216_0.jpg CHANGED

Git LFS Details

  • SHA256: c4a5ef66a593d75d396664ba0d5599863d327611e5bcadf94583aca7aa3b1583
  • Pointer size: 131 Bytes
  • Size of remote file: 630 kB

Git LFS Details

  • SHA256: 2b4462c8cb70cd361ea237bc80c4e0e2bb79d74f3eb4d9f8b5d56c7ab5dcb8f1
  • Pointer size: 131 Bytes
  • Size of remote file: 183 kB
samples/unet_1280x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: 4ed6add485367036b8768bc035d6eda3345a8fdeb1b0b957bb0648949dad71c4
  • Pointer size: 131 Bytes
  • Size of remote file: 514 kB

Git LFS Details

  • SHA256: a7263afe1067cf6b20e31c677c1944e20c31293ed2e57b2a9d6b01efae057419
  • Pointer size: 131 Bytes
  • Size of remote file: 740 kB
samples/unet_1280x640_0.jpg CHANGED

Git LFS Details

  • SHA256: f178b1814a400e1ebd6cf5a7ecf55338ad0464c917aac116c19ddc796000c791
  • Pointer size: 131 Bytes
  • Size of remote file: 410 kB

Git LFS Details

  • SHA256: 61ca5d03ab929e5713d342d5a2f1e19674f8b78f34ad7dda8a35306fd532bb1d
  • Pointer size: 131 Bytes
  • Size of remote file: 339 kB
samples/unet_1280x704_0.jpg CHANGED

Git LFS Details

  • SHA256: cfa560ec714906e5be7d1649a12d8fd1ae5e9bb4316fb53e83a27721470b8f9e
  • Pointer size: 131 Bytes
  • Size of remote file: 254 kB

Git LFS Details

  • SHA256: 95eb3b1db0c31d0cd17d3f41efc8bbc06dc1310b82e16f62512c80fb904136a4
  • Pointer size: 131 Bytes
  • Size of remote file: 325 kB
samples/unet_1280x768_0.jpg CHANGED

Git LFS Details

  • SHA256: 0bf20aef4b65e34cf2b3bd04eb5aa0163b9eff407e8ff0949fc8037ed1e4c00f
  • Pointer size: 131 Bytes
  • Size of remote file: 145 kB

Git LFS Details

  • SHA256: 648bb05056c218bf0e2d6ec51b8cddf0e0a326a47340be508494c2a39d0c870a
  • Pointer size: 131 Bytes
  • Size of remote file: 235 kB
samples/unet_1280x832_0.jpg CHANGED

Git LFS Details

  • SHA256: a8a3bf972ec1ffa151103716ff42fe9c3ef4ce06e576778a8ba93adc7fe4db5e
  • Pointer size: 131 Bytes
  • Size of remote file: 367 kB

Git LFS Details

  • SHA256: 1dc69a1efc03ac0c9f28875bb01702ac17290a3d6f258aaba1122f110f078d29
  • Pointer size: 131 Bytes
  • Size of remote file: 255 kB
samples/unet_1280x896_0.jpg CHANGED

Git LFS Details

  • SHA256: d70e21c44e5f1377037a353c377b2baf9a1fb0f50521267e60e9c2bb5789e4cf
  • Pointer size: 131 Bytes
  • Size of remote file: 315 kB

Git LFS Details

  • SHA256: 166d377cb79b576561038254ee66dfcc48870644f8c11ad668115f7f81ba0516
  • Pointer size: 131 Bytes
  • Size of remote file: 441 kB
samples/unet_1280x960_0.jpg CHANGED

Git LFS Details

  • SHA256: a2a706225020bb74744b1401821f458fb35d69eab8b05fa2a296223935eda151
  • Pointer size: 131 Bytes
  • Size of remote file: 274 kB

Git LFS Details

  • SHA256: 05d850414b00cce7d015c2d07915d29635e102376d3e978a40d63a2586364d82
  • Pointer size: 131 Bytes
  • Size of remote file: 213 kB
samples/unet_640x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: e4c3552e45ba48b5d62087f17bc73367b2d5184e2bf188c3c83cccb446ded886
  • Pointer size: 131 Bytes
  • Size of remote file: 104 kB

Git LFS Details

  • SHA256: 64875ab6d6a9154334b1aced2bfafb1bc353a96f31e98d0917a1a712b50daab3
  • Pointer size: 131 Bytes
  • Size of remote file: 138 kB
samples/unet_704x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: 38e9fc74c1fb6610524dfc4d7761f14e2e109139d61f900507c7a2e29e7e0d78
  • Pointer size: 131 Bytes
  • Size of remote file: 415 kB

Git LFS Details

  • SHA256: b87d04d5ce6c1dc77cb6a81f0267898e1eaf7a6f0b0830c68b40ea2bd2eb40b7
  • Pointer size: 131 Bytes
  • Size of remote file: 299 kB
samples/unet_768x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: 35273b2079dbc26c81b7f9c31ba4bb66876b615a3c1f6a0319a3085efd56b00c
  • Pointer size: 131 Bytes
  • Size of remote file: 384 kB

Git LFS Details

  • SHA256: 70fd47235f849f43b24f319922820d0a8c8c52c44214c36db4f1908355234c86
  • Pointer size: 131 Bytes
  • Size of remote file: 506 kB
samples/unet_832x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: f548178489b12320aab6e4481a9b6d86ffb98ce73b45d847475c92ec61943a8d
  • Pointer size: 131 Bytes
  • Size of remote file: 206 kB

Git LFS Details

  • SHA256: 684cfc7a6f930492a973a1fabb6c90fd03894548b3184a4ecd18a3771573e54d
  • Pointer size: 131 Bytes
  • Size of remote file: 444 kB
samples/unet_896x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: df7e07060977bf296389f8975b90f6303e82cb1c25a3b279d48689ffaa67780d
  • Pointer size: 131 Bytes
  • Size of remote file: 207 kB

Git LFS Details

  • SHA256: 7fe3ead118cd85b61c012405cce867e11c73ab7492e6b267ba6555f42fa1a962
  • Pointer size: 131 Bytes
  • Size of remote file: 241 kB
samples/unet_960x1280_0.jpg CHANGED

Git LFS Details

  • SHA256: 248fa863a456e1c40c6c23090505b532481b03d9679f1be58bacabcd5319338c
  • Pointer size: 131 Bytes
  • Size of remote file: 312 kB

Git LFS Details

  • SHA256: df6610b355ebd82a9f1a779dddcdb5836e53626afaf0f8db33c88a0a3140383e
  • Pointer size: 131 Bytes
  • Size of remote file: 545 kB
test.ipynb CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:faf7be1e5cdebb1b6183b35157a275dff605c4a03ade67a19da702b683b658c7
3
- size 5815496
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:57663fad435575b07ff8ce6a7866cc47fc1c747eec2428030b8f7648304f65be
3
+ size 5834915
train-Copy1.py ADDED
@@ -0,0 +1,917 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from comet_ml import Experiment
2
+ import os
3
+ os.environ["NCCL_P2P_DISABLE"] = "1"
4
+ os.environ["NCCL_IB_DISABLE"] = "1" # test it
5
+ os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
6
+ import math
7
+ import torch
8
+ import numpy as np
9
+ import matplotlib.pyplot as plt
10
+ from torch.utils.data import DataLoader, Sampler
11
+ from torch.utils.data.distributed import DistributedSampler
12
+ from torch.optim.lr_scheduler import LambdaLR
13
+ from collections import defaultdict
14
+ from diffusers import UNet2DConditionModel,AutoencoderKL,FlowMatchEulerDiscreteScheduler
15
+ from accelerate import Accelerator, DeepSpeedPlugin
16
+ from datasets import load_from_disk
17
+ from tqdm import tqdm
18
+ from PIL import Image, ImageOps
19
+ import wandb
20
+ import random,time
21
+ import gc
22
+ from accelerate.state import DistributedType
23
+ from torch.distributed import broadcast_object_list
24
+ from torch.utils.checkpoint import checkpoint
25
+ from diffusers.models.attention_processor import AttnProcessor2_0
26
+ from datetime import datetime
27
+ import bitsandbytes as bnb
28
+ import torch.nn.functional as F
29
+ from collections import deque
30
+ from transformers import Qwen3_5Tokenizer, Qwen3_5ForConditionalGeneration
31
+ import argparse
32
+
33
+ # --------------------------- Параметры ---------------------------
34
+ ds_path = "datasets/ds1234_1280"
35
+ project = "unet"
36
+ # 1. Считаем локальный батч для ОДНОЙ карты (3 на каждые 32 Гб)
37
+ gpu_mem_gb = torch.cuda.get_device_properties(0).total_memory / 1e9
38
+ local_bs = max(1, int((gpu_mem_gb / 32) * 2))
39
+ # 2. Умножаем на количество ГПУ, чтобы получить ГЛОБАЛЬНЫЙ батч
40
+ num_gpus = torch.cuda.device_count()
41
+ ## total batch (split // num `GPU)
42
+ batch_size = local_bs * num_gpus
43
+ print(f"GPUs: {num_gpus}, Local BS: {local_bs}, Global BS: {local_bs * num_gpus}")
44
+ base_learning_rate = 2e-5
45
+ min_learning_rate = 3e-6
46
+ num_epochs = num_gpus #8 * max(1, int(num_gpus / 2))
47
+ sample_interval_share = 20
48
+ cfg_dropout = 0.10
49
+ max_length = 248
50
+ use_wandb = False
51
+ use_comet_ml = True
52
+ save_model = True
53
+ use_decay = True
54
+ fbp = False
55
+ optimizer_type = "adam8bit"
56
+ torch_compile = False
57
+ unet_gradient = True
58
+ loss_normalize = False
59
+ fixed_seed = False
60
+ shuffle = True
61
+ comet_ml_api_key = "Agctp26mbqnoYrrlvQuKSTk6r"
62
+ comet_ml_workspace = "recoilme"
63
+ torch.backends.cuda.matmul.allow_tf32 = True
64
+ torch.backends.cudnn.allow_tf32 = True
65
+ # Включение Flash Attention 2/SDPA #MAX_JOBS=4 pip install flash-attn --no-build-isolation
66
+ torch.backends.cuda.enable_flash_sdp(True)
67
+ torch.backends.cuda.enable_mem_efficient_sdp(True)
68
+ torch.backends.cuda.enable_math_sdp(False) # Отключаем медленный вариант
69
+ save_barrier = 1.25
70
+ warmup_percent = 0.01
71
+ #percentile_clipping = 95
72
+ betta2 = 0.995
73
+ eps = 1e-7
74
+ clip_grad_norm = 1.0
75
+ limit = 0
76
+ checkpoints_folder = ""
77
+ gradient_accumulation_steps = 1
78
+ dtype = torch.float32
79
+ mixed_precision = "no"
80
+
81
+ # Параметры для диффузии
82
+ n_diffusion_steps = 40
83
+ samples_to_generate = 12
84
+ guidance_scale = 4
85
+
86
+ # Папки для сохранения результатов
87
+ generated_folder = "samples"
88
+ os.makedirs(generated_folder, exist_ok=True)
89
+
90
+ # Настройка seed
91
+ current_date = datetime.now()
92
+ seed = int(current_date.strftime("%Y%m%d")) + 42
93
+ if fixed_seed:
94
+ torch.manual_seed(seed)
95
+ np.random.seed(seed)
96
+ random.seed(seed)
97
+ if torch.cuda.is_available():
98
+ torch.cuda.manual_seed_all(seed)
99
+
100
+ accelerator = Accelerator(
101
+ mixed_precision=mixed_precision,
102
+ gradient_accumulation_steps=gradient_accumulation_steps
103
+ )
104
+ device = accelerator.device
105
+
106
+ print("init")
107
+ # Создаём объект ArgumentParser с рассчитанными значениями по умолчанию
108
+ parser = argparse.ArgumentParser(description='Train a model on a dataset.')
109
+ parser.add_argument('--ds-path', type=str, default=ds_path, help='Path to the dataset')
110
+ parser.add_argument('--ep', type=int, default=num_epochs, help='Number of epochs to train the model')
111
+ parser.add_argument('--batch', type=int, default=batch_size, help='Total batch size')
112
+ parser.add_argument('--min-lr', type=float, default=min_learning_rate, help='Minimum learning rate')
113
+ parser.add_argument('--max-lr', type=float, default=base_learning_rate, help='Maximum learning rate')
114
+ parser.add_argument('--dry-run', action='store_true',default=False, help='Run configuration without saving/sampling')
115
+
116
+ # Парсим аргументы командной строки
117
+ args = parser.parse_args()
118
+
119
+ # Используем значения из аргументов
120
+ ds_path = args.ds_path
121
+ base_learning_rate = args.max_lr
122
+ min_learning_rate = args.min_lr
123
+ num_epochs = args.ep
124
+ if args.dry_run:
125
+ save_model = False
126
+
127
+ # --------------------------- Инициализация WandB ---------------------------
128
+ if accelerator.is_main_process:
129
+ if use_wandb:
130
+ wandb.init(project=project, config={
131
+ "batch_size": batch_size,
132
+ "base_learning_rate": base_learning_rate,
133
+ "num_epochs": num_epochs,
134
+ "optimizer_type": optimizer_type,
135
+ })
136
+ if use_comet_ml:
137
+ from comet_ml import Experiment
138
+ comet_experiment = Experiment(
139
+ api_key=comet_ml_api_key,
140
+ project_name=project,
141
+ workspace=comet_ml_workspace
142
+ )
143
+ hyper_params = {
144
+ "batch_size": batch_size,
145
+ "base_learning_rate": base_learning_rate,
146
+ "num_epochs": num_epochs,
147
+ }
148
+ comet_experiment.log_parameters(hyper_params)
149
+
150
+ # --------------------------- Загрузка моделей ---------------------------
151
+ vae = AutoencoderKL.from_pretrained("vae", torch_dtype=dtype).to(device).eval()
152
+ tokenizer = Qwen3_5Tokenizer.from_pretrained("tokenizer")
153
+ text_encoder = Qwen3_5ForConditionalGeneration.from_pretrained("text_encoder", torch_dtype=torch.float16).to(device).eval()
154
+ scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained("scheduler")
155
+
156
+ def encode_texts(texts, max_length=max_length):
157
+ if texts is None:
158
+ texts = [""]
159
+ if isinstance(texts, str):
160
+ texts = [texts]
161
+
162
+ with torch.no_grad():
163
+
164
+ # --- 2. QWEN Энкодер (через Chat Template) ---
165
+ # 1. Собираем текстовые промпты оборачивая их в Chat Template
166
+ formatted_prompts = []
167
+ for t in texts:
168
+ messages = [{"role": "user", "content": [{"type": "text", "text": t}]}]
169
+ res_text = tokenizer.apply_chat_template(
170
+ messages,
171
+ add_generation_prompt=True,
172
+ tokenize=False
173
+ )
174
+ formatted_prompts.append(res_text)
175
+
176
+ # 2. Токенизируем, режем и добавляем паддинг за один раз
177
+ toks = tokenizer(
178
+ formatted_prompts,
179
+ padding="max_length",
180
+ max_length=max_length,
181
+ truncation=True,
182
+ return_tensors="pt"
183
+ ).to(device)
184
+
185
+ # 3. Прогоняем через модель
186
+ outputs = text_encoder(
187
+ input_ids=toks.input_ids,
188
+ attention_mask=toks.attention_mask,
189
+ output_hidden_states=True
190
+ )
191
+
192
+ layer_index = -2
193
+ last_hidden = outputs.hidden_states[layer_index]
194
+ seq_len = toks.attention_mask.sum(dim=1) - 1
195
+ pooled = last_hidden[torch.arange(len(last_hidden)), seq_len.clamp(min=0)]
196
+ #pooled = torch.cat([pooled_clip, pooled], dim=1)
197
+ return last_hidden.to(dtype), toks.attention_mask, pooled.to(dtype)
198
+
199
+ shift_factor = getattr(vae.config, "shift_factor", 0.0)
200
+ if shift_factor is None:
201
+ shift_factor = 0.0
202
+
203
+ scaling_factor = getattr(vae.config, "scaling_factor", 1.0)
204
+ if scaling_factor is None:
205
+ scaling_factor = 1.0
206
+
207
+ mean = getattr(vae.config, "latents_mean", None)
208
+ std = getattr(vae.config, "latents_std", None)
209
+ if mean is not None and std is not None:
210
+ latents_std = torch.tensor(std, device=device, dtype=dtype).view(1, len(std), 1, 1)
211
+ latents_mean = torch.tensor(mean, device=device, dtype=dtype).view(1, len(mean), 1, 1)
212
+
213
+ import numpy as np
214
+ from torch.utils.data import Sampler
215
+
216
+
217
+ class DistributedResolutionBatchSampler(Sampler):
218
+ def __init__(self, dataset, batch_size, num_replicas, rank, drop_last=True, shuffle=True):
219
+ self.dataset = dataset
220
+ self.num_replicas = num_replicas
221
+ self.rank = rank
222
+ self.shuffle = shuffle
223
+ self.drop_last = drop_last
224
+ self.epoch = 0
225
+
226
+ # batch на одну GPU
227
+ self.batch_size = max(1, batch_size // num_replicas)
228
+ self.global_batch = self.batch_size * num_replicas
229
+
230
+ try:
231
+ widths = np.asarray(dataset["width"])
232
+ heights = np.asarray(dataset["height"])
233
+ except KeyError:
234
+ widths = np.zeros(len(dataset))
235
+ heights = np.zeros(len(dataset))
236
+
237
+ # --- группировка индексов ---
238
+ groups = {}
239
+ for i, (w, h) in enumerate(zip(widths, heights)):
240
+ groups.setdefault((w, h), []).append(i)
241
+
242
+ # --- создаём список всех глобальных батчей ---
243
+ all_batches = []
244
+
245
+ for indices in groups.values():
246
+
247
+ idx = np.asarray(indices, dtype=np.int64)
248
+
249
+ num_batches = len(idx) // self.global_batch
250
+ if num_batches == 0:
251
+ continue
252
+
253
+ idx = idx[: num_batches * self.global_batch]
254
+
255
+ batches = idx.reshape(num_batches, self.global_batch)
256
+
257
+ all_batches.append(batches)
258
+
259
+ if len(all_batches) > 0:
260
+ self.global_batches = np.concatenate(all_batches, axis=0)
261
+ else:
262
+ self.global_batches = np.empty((0, self.global_batch), dtype=np.int64)
263
+
264
+ self.num_batches = len(self.global_batches)
265
+
266
+ def __iter__(self):
267
+
268
+ rng = np.random.RandomState(self.epoch)
269
+
270
+ order = np.arange(self.num_batches)
271
+
272
+ if self.shuffle:
273
+ rng.shuffle(order)
274
+
275
+ start = self.rank * self.batch_size
276
+ end = start + self.batch_size
277
+
278
+ for i in order:
279
+ yield self.global_batches[i][start:end]
280
+
281
+ def __len__(self):
282
+ return self.num_batches
283
+
284
+ def set_epoch(self, epoch):
285
+ self.epoch = epoch
286
+
287
+ class DistributedResolutionBatchSamplerOld(Sampler):
288
+ def __init__(self, dataset, batch_size, num_replicas, rank, drop_last=True, shuffle=False):
289
+ self.dataset = dataset
290
+ self.num_replicas = num_replicas
291
+ self.rank = rank
292
+ self.shuffle = shuffle
293
+ self.drop_last = drop_last
294
+ self.epoch = 0
295
+
296
+ # batch на одну GPU
297
+ self.batch_size = max(1, batch_size // num_replicas)
298
+ self.global_batch = self.batch_size * num_replicas
299
+
300
+ try:
301
+ widths = np.asarray(dataset["width"])
302
+ heights = np.asarray(dataset["height"])
303
+ except KeyError:
304
+ widths = np.zeros(len(dataset))
305
+ heights = np.zeros(len(dataset))
306
+
307
+ # --- группировка индексов ---
308
+ groups = {}
309
+ for i, (w, h) in enumerate(zip(widths, heights)):
310
+ groups.setdefault((w, h), []).append(i)
311
+
312
+ # --- строим батчи один раз (кеш) ---
313
+ self.group_batches = []
314
+
315
+ for indices in groups.values():
316
+
317
+ idx = np.asarray(indices, dtype=np.int64)
318
+
319
+ num_batches = len(idx) // self.global_batch
320
+ if num_batches == 0:
321
+ continue
322
+
323
+ idx = idx[: num_batches * self.global_batch]
324
+
325
+ batches = idx.reshape(num_batches, self.global_batch)
326
+
327
+ self.group_batches.append(batches)
328
+
329
+ # число батчей
330
+ self.num_batches = sum(len(g) for g in self.group_batches)
331
+
332
+ def __iter__(self):
333
+
334
+ rng = np.random.RandomState(self.epoch)
335
+
336
+ groups = []
337
+
338
+ # shuffle внутри групп
339
+ for g in self.group_batches:
340
+
341
+ order = np.arange(len(g))
342
+
343
+ if self.shuffle:
344
+ rng.shuffle(order)
345
+
346
+ groups.append(g[order])
347
+
348
+ # shuffle порядок групп
349
+ if self.shuffle:
350
+ rng.shuffle(groups)
351
+
352
+ # --- round robin сборка ---
353
+ group_pos = [0] * len(groups)
354
+
355
+ start = self.rank * self.batch_size
356
+ end = start + self.batch_size
357
+
358
+ remaining = True
359
+
360
+ while remaining:
361
+
362
+ remaining = False
363
+
364
+ for gi, g in enumerate(groups):
365
+
366
+ pos = group_pos[gi]
367
+
368
+ if pos < len(g):
369
+
370
+ batch = g[pos]
371
+
372
+ group_pos[gi] += 1
373
+ remaining = True
374
+
375
+ yield batch[start:end]
376
+
377
+ def __len__(self):
378
+ return self.num_batches
379
+
380
+ def set_epoch(self, epoch):
381
+ self.epoch = epoch
382
+
383
+
384
+ # --- [UPDATED] Функция для фиксированных семплов ---
385
+ def get_fixed_samples_by_resolution(dataset, samples_per_group=1):
386
+ size_groups = defaultdict(list)
387
+ try:
388
+ widths = dataset["width"]
389
+ heights = dataset["height"]
390
+ except KeyError:
391
+ widths = [0] * len(dataset)
392
+ heights = [0] * len(dataset)
393
+ for i, (w, h) in enumerate(zip(widths, heights)):
394
+ size = (w, h)
395
+ size_groups[size].append(i)
396
+
397
+ fixed_samples = {}
398
+ for size, indices in size_groups.items():
399
+ n_samples = min(samples_per_group, len(indices))
400
+ if len(size_groups)==1:
401
+ n_samples = samples_to_generate
402
+ if n_samples == 0:
403
+ continue
404
+ sample_indices = random.sample(indices, n_samples)
405
+ samples_data = [dataset[idx] for idx in sample_indices]
406
+
407
+ latents = torch.tensor(np.array([item["vae"] for item in samples_data])).to(device=device, dtype=dtype)
408
+ texts = [item["text"] for item in samples_data]
409
+
410
+ # Кодируем тексты на лету, чтобы получить маски и пулинг
411
+ embeddings, masks, pooled = encode_texts(texts)
412
+
413
+ fixed_samples[size] = (latents, embeddings, masks, texts, pooled)
414
+
415
+ print(f"Создано {len(fixed_samples)} групп фиксированных семплов по разрешениям")
416
+ return fixed_samples
417
+
418
+ if limit > 0:
419
+ dataset = load_from_disk(ds_path).select(range(limit))
420
+ else:
421
+ dataset = load_from_disk(ds_path)
422
+
423
+ dataset = dataset.filter(
424
+ lambda x: [not (path.startswith("//workspace/ds/animesfw") or path.startswith("//workspace/animesfw") or path.startswith("//workspace/ds/d4/animesfw")) for path in x["image_path"]],
425
+
426
+ batched=True,
427
+ batch_size=10000, # обрабатываем по 10к строк за раз
428
+ num_proc=8
429
+ )
430
+ print(f"Осталось примеров после фильтрации: {len(dataset)}")
431
+
432
+ # --- Collate Function ---
433
+ def collate_fn_simple(batch):
434
+ # 1. Латенты (VAE)
435
+ #latents = torch.tensor(np.array([item["vae"] for item in batch])).to(device, dtype=dtype)
436
+ latents = torch.from_numpy(np.array([item["vae"] for item in batch], dtype=np.float16)).to(device, dtype=dtype)
437
+
438
+ # 2. Текст берем сырой из датасета
439
+ raw_texts = [item["text"] for item in batch]
440
+ texts = [
441
+ "" if t.lower().startswith("zero")
442
+ else "" if random.random() < cfg_dropout
443
+ else t[1:].lstrip() if t.startswith(".")
444
+ else t.replace("The image shows ", "").replace("The image is ", "").replace("This image captures ","").strip()
445
+ for t in raw_texts
446
+ ]
447
+ # 3. Кодируем на лету
448
+ # Возвращает: hidden (B, L, D), mask (B, L)
449
+ embeddings, attention_mask, pooled = encode_texts(texts)
450
+
451
+ # attention_mask от токенизатора уже имеет нужный формат, но на всякий случай приведем к long
452
+ attention_mask = attention_mask.to(dtype=torch.int64)
453
+
454
+ return latents, embeddings, attention_mask, pooled
455
+
456
+ batch_sampler = DistributedResolutionBatchSampler(
457
+ dataset=dataset,
458
+ batch_size=batch_size,
459
+ num_replicas=accelerator.num_processes,
460
+ rank=accelerator.process_index,
461
+ shuffle = shuffle
462
+ )
463
+
464
+ dataloader = DataLoader(dataset, batch_sampler=batch_sampler, collate_fn=collate_fn_simple)
465
+
466
+ if accelerator.is_main_process:
467
+ print("Total samples", len(dataloader))
468
+ dataloader = accelerator.prepare(dataloader)
469
+
470
+ start_epoch = 0
471
+ global_step = 0
472
+ total_training_steps = (len(dataloader) * num_epochs)
473
+ world_size = accelerator.state.num_processes
474
+
475
+ # Загрузка UNet
476
+ latest_checkpoint = os.path.join(checkpoints_folder, project)
477
+ if os.path.isdir(latest_checkpoint):
478
+ print("Загружаем UNet из чекпоинта:", latest_checkpoint)
479
+ unet = UNet2DConditionModel.from_pretrained(latest_checkpoint).to(device=device, dtype=dtype)
480
+ if unet_gradient:
481
+ unet.enable_gradient_checkpointing()
482
+ unet.set_use_memory_efficient_attention_xformers(False)
483
+ try:
484
+ unet.set_attn_processor(AttnProcessor2_0())
485
+ except Exception as e:
486
+ print(f"Ошибка при включении SDPA: {e}")
487
+ unet.set_use_memory_efficient_attention_xformers(True)
488
+ else:
489
+ raise FileNotFoundError(f"UNet checkpoint not found at {latest_checkpoint}")
490
+
491
+
492
+ def create_optimizer(name, params):
493
+ if name == "adam8bit":
494
+ return bnb.optim.AdamW8bit(
495
+ params, lr=base_learning_rate, betas=(0.9, betta2), eps=eps, weight_decay=0.01,
496
+ #percentile_clipping=percentile_clipping
497
+ )
498
+ elif name == "adam":
499
+ return torch.optim.AdamW(
500
+ params, lr=base_learning_rate, betas=(0.9, betta2), eps=1e-8, weight_decay=0.01
501
+ )
502
+ else:
503
+ raise ValueError(f"Unknown optimizer: {name}")
504
+
505
+ if fbp:
506
+ trainable_params = list(unet.parameters())
507
+ optimizer_dict = {p: create_optimizer(optimizer_type, [p]) for p in trainable_params}
508
+ def optimizer_hook(param):
509
+ optimizer_dict[param].step()
510
+ optimizer_dict[param].zero_grad(set_to_none=True)
511
+ for param in trainable_params:
512
+ param.register_post_accumulate_grad_hook(optimizer_hook)
513
+ unet, optimizer = accelerator.prepare(unet, optimizer_dict)
514
+ else:
515
+ # 1. Сначала замораживаем ВСЕ параметры UNet
516
+ #unet.requires_grad_(False)
517
+
518
+ # 2. Размораживаем только нужные
519
+ #trainable_params_names = ["conv_in.weight", "conv_in.bias", "conv_out.weight", "conv_out.bias"]
520
+ #train_params = []
521
+
522
+ #for name, param in unet.named_parameters():
523
+ # if any(target in name for target in trainable_params_names):
524
+ # param.requires_grad = True
525
+ # train_params.append(param)
526
+ # print(f"Обучаемый слой: {name}")
527
+
528
+ unet.requires_grad_(True)
529
+ optimizer = create_optimizer(optimizer_type, unet.parameters())
530
+
531
+ def lr_schedule(step):
532
+ x = step / (total_training_steps * world_size)
533
+ warmup = warmup_percent
534
+ if not use_decay:
535
+ return base_learning_rate
536
+ if x < warmup:
537
+ return min_learning_rate + (base_learning_rate - min_learning_rate) * (x / warmup)
538
+ decay_ratio = (x - warmup) / (1 - warmup)
539
+ return min_learning_rate + 0.5 * (base_learning_rate - min_learning_rate) * \
540
+ (1 + math.cos(math.pi * decay_ratio))
541
+ lr_scheduler = LambdaLR(optimizer, lambda step: lr_schedule(step) / base_learning_rate)
542
+ unet, optimizer, lr_scheduler = accelerator.prepare(unet, optimizer, lr_scheduler)
543
+
544
+ if torch_compile:
545
+ print("compiling")
546
+ unet = torch.compile(unet)
547
+ print("compiling - ok")
548
+
549
+ # Фиксированные семплы
550
+ fixed_samples = get_fixed_samples_by_resolution(dataset)
551
+
552
+ # --- [UPDATED] Функция для негативного эмбеддинга (возвращает 3 элемента) ---
553
+ def get_negative_embedding(neg_prompt="", batch_size=1):
554
+ if not neg_prompt:
555
+ hidden_dim = 2048
556
+ seq_len = max_length
557
+ empty_emb = torch.zeros((batch_size, seq_len, hidden_dim), dtype=dtype, device=device)
558
+ empty_mask = torch.ones((batch_size, seq_len), dtype=torch.int64, device=device)
559
+ return empty_emb, empty_mask
560
+
561
+ uncond_emb, uncond_mask, uncond_pooled = encode_texts([neg_prompt])
562
+ uncond_emb = uncond_emb.to(dtype=dtype, device=device).repeat(batch_size, 1, 1)
563
+ uncond_mask = uncond_mask.to(device=device).repeat(batch_size, 1)
564
+ uncond_pooled = uncond_pooled.to(device=device).repeat(batch_size, 1)
565
+
566
+ return uncond_emb, uncond_mask, uncond_pooled
567
+
568
+ # Получаем негативные (пустые) условия для валидации
569
+ uncond_emb, uncond_mask, uncond_pooled = get_negative_embedding("low quality")
570
+
571
+ # --- Функция генерации семплов ---
572
+ @torch.compiler.disable()
573
+ @torch.no_grad()
574
+ def generate_and_save_samples(fixed_samples_cpu, uncond_data, step):
575
+ uncond_emb, uncond_mask, uncond_pooled = uncond_data
576
+
577
+ original_model = None
578
+ try:
579
+ if not torch_compile:
580
+ original_model = accelerator.unwrap_model(unet, keep_torch_compile=True).eval()
581
+ else:
582
+ original_model = unet.eval()
583
+
584
+ vae.to(device=device).eval()
585
+
586
+ all_generated_images = []
587
+ all_captions = []
588
+
589
+ # Распаковываем 5 элементов (добавились mask)
590
+ for size, (sample_latents, sample_text_embeddings, sample_mask, sample_text, sample_pooled) in fixed_samples_cpu.items():
591
+ width, height = size
592
+ sample_latents = sample_latents.to(dtype=dtype, device=device)
593
+ sample_text_embeddings = sample_text_embeddings.to(dtype=dtype, device=device)
594
+ sample_mask = sample_mask.to(device=device)
595
+ sample_pooled = sample_pooled.to(dtype=dtype, device=device)
596
+
597
+ latents = torch.randn(
598
+ sample_latents.shape,
599
+ device=device,
600
+ dtype=sample_latents.dtype,
601
+ generator=torch.Generator(device=device).manual_seed(seed)
602
+ )
603
+
604
+ scheduler.set_timesteps(n_diffusion_steps, device=device)
605
+
606
+ time_ids = torch.zeros(
607
+ sample_pooled.shape[0], # ← вот это главное
608
+ 6,
609
+ device=device,
610
+ dtype=torch.long
611
+ )
612
+
613
+ for t in scheduler.timesteps:
614
+ if guidance_scale != 1:
615
+ latent_model_input = torch.cat([latents, latents], dim=0)
616
+
617
+ curr_batch_size = sample_text_embeddings.shape[0]
618
+ seq_len = sample_text_embeddings.shape[1]
619
+ hidden_dim = sample_text_embeddings.shape[2]
620
+
621
+ neg_emb_batch = uncond_emb[0:1].expand(curr_batch_size, -1, -1)
622
+ text_embeddings_batch = torch.cat([neg_emb_batch, sample_text_embeddings], dim=0)
623
+
624
+ neg_mask_batch = uncond_mask[0:1].expand(curr_batch_size, -1)
625
+ attention_mask_batch = torch.cat([neg_mask_batch, sample_mask], dim=0)
626
+
627
+ neg_pooled_batch = uncond_pooled[0:1].expand(curr_batch_size, -1)
628
+ pooled_batch = torch.cat([neg_pooled_batch, sample_pooled], dim=0)
629
+
630
+ # ← КЛЮЧЕВОЕ ИСПРАВЛЕНИЕ — time_ids под текущий удвоенный батч!
631
+ time_ids = torch.zeros(
632
+ pooled_batch.shape[0], # 2 * curr_batch_size при CFG
633
+ 6,
634
+ device=device,
635
+ dtype=torch.long
636
+ )
637
+
638
+ else:
639
+ latent_model_input = latents
640
+ text_embeddings_batch = sample_text_embeddings
641
+ attention_mask_batch = sample_mask
642
+ pooled_batch = sample_pooled
643
+
644
+ time_ids = torch.zeros(
645
+ pooled_batch.shape[0],
646
+ 6,
647
+ device=device,
648
+ dtype=torch.long
649
+ )
650
+
651
+ # Теперь всё имеет одинаковый batch size
652
+ model_out = original_model(
653
+ latent_model_input,
654
+ t,
655
+ encoder_hidden_states=text_embeddings_batch,
656
+ encoder_attention_mask=attention_mask_batch,
657
+ added_cond_kwargs={
658
+ "text_embeds": pooled_batch,
659
+ "time_ids": time_ids
660
+ },
661
+ )
662
+
663
+ flow = getattr(model_out, "sample", model_out)
664
+
665
+ if guidance_scale != 1:
666
+ flow_uncond, flow_cond = flow.chunk(2)
667
+ flow = flow_uncond + guidance_scale * (flow_cond - flow_uncond)
668
+
669
+ latents = scheduler.step(flow, t, latents).prev_sample
670
+
671
+ current_latents = latents
672
+ if step==0:
673
+ current_latents = sample_latents
674
+
675
+ if latents_mean is not None and latents_std is not None:
676
+ latents = current_latents * latents_std + latents_mean
677
+
678
+ decoded = vae.decode(latents.to(torch.float32)).sample
679
+ decoded_fp32 = decoded.to(torch.float32)
680
+
681
+ for img_idx, img_tensor in enumerate(decoded_fp32):
682
+ img = (img_tensor / 2 + 0.5).clamp(0, 1).cpu().numpy()
683
+ img = img.transpose(1, 2, 0)
684
+
685
+ if np.isnan(img).any():
686
+ print("NaNs found, saving stopped! Step:", step)
687
+ pil_img = Image.fromarray((img * 255).astype("uint8"))
688
+
689
+ max_w_overall = max(s[0] for s in fixed_samples_cpu.keys())
690
+ max_h_overall = max(s[1] for s in fixed_samples_cpu.keys())
691
+ max_w_overall = max(255, max_w_overall)
692
+ max_h_overall = max(255, max_h_overall)
693
+
694
+ padded_img = ImageOps.pad(pil_img, (max_w_overall, max_h_overall), color='white')
695
+ all_generated_images.append(padded_img)
696
+
697
+ caption_text = sample_text[img_idx][:300] if img_idx < len(sample_text) else ""
698
+ all_captions.append(caption_text)
699
+
700
+ sample_path = f"{generated_folder}/{project}_{width}x{height}_{img_idx}.jpg"
701
+ pil_img.save(sample_path, "JPEG", quality=95)
702
+
703
+ if use_wandb and accelerator.is_main_process:
704
+ wandb_images = [
705
+ wandb.Image(img, caption=f"{all_captions[i]}")
706
+ for i, img in enumerate(all_generated_images)
707
+ ]
708
+ wandb.log({"generated_images": wandb_images})
709
+ if use_comet_ml and accelerator.is_main_process:
710
+ for i, img in enumerate(all_generated_images):
711
+ comet_experiment.log_image(
712
+ image_data=img,
713
+ name=f"step_{step}_img_{i}",
714
+ step=step,
715
+ metadata={"caption": all_captions[i]}
716
+ )
717
+ finally:
718
+ vae.to("cpu")
719
+ try:
720
+ all_generated_images.clear()
721
+ all_captions.clear()
722
+ del all_generated_images, all_captions
723
+ del latents, current_latents, latent_model_input, flow
724
+ del decoded, decoded_fp32
725
+ del sample_latents, sample_text_embeddings, sample_mask, sample_pooled # Копии на GPU
726
+ del model_out
727
+ except UnboundLocalError:
728
+ pass
729
+
730
+ # 3. Синхронизируем CUDA перед очисткой
731
+ torch.cuda.synchronize()
732
+ # 4. Теперь чистим кэш аллокатора и вызываем GC
733
+ torch.cuda.empty_cache()
734
+ gc.collect()
735
+
736
+ # --------------------------- Генерация сэмплов перед обучением ---------------------------
737
+ if accelerator.is_main_process:
738
+ if save_model:
739
+ print("Генерация сэмплов до старта обучения...")
740
+ generate_and_save_samples(fixed_samples, (uncond_emb, uncond_mask,uncond_pooled), 0)
741
+ accelerator.wait_for_everyone()
742
+
743
+ def save_checkpoint(unet, variant=""):
744
+ if accelerator.is_main_process:
745
+ model_to_save = None
746
+ if not torch_compile:
747
+ model_to_save = accelerator.unwrap_model(unet)
748
+ else:
749
+ model_to_save = unet
750
+
751
+ if variant != "":
752
+ model_to_save.to(dtype=torch.float16).save_pretrained(
753
+ os.path.join(checkpoints_folder, f"{project}"), variant=variant
754
+ )
755
+ else:
756
+ model_to_save.save_pretrained(os.path.join(checkpoints_folder, f"{project}"))
757
+
758
+ torch.cuda.synchronize()
759
+ torch.cuda.empty_cache()
760
+ gc.collect()
761
+ #unet = unet.to(dtype=dtype) #TODO: wtf???
762
+
763
+ # --------------------------- Тренировочный цикл ---------------------------
764
+ if accelerator.is_main_process:
765
+ print(f"Total steps per GPU: {total_training_steps}")
766
+
767
+ epoch_loss_points = []
768
+ progress_bar = tqdm(total=total_training_steps, disable=not accelerator.is_local_main_process, desc="Training", unit="step")
769
+
770
+ steps_per_epoch = len(dataloader)
771
+ sample_interval = max(1, steps_per_epoch // sample_interval_share)
772
+ min_loss = 4.
773
+ last_sample_time = time.time()
774
+ sample_interval_seconds = 60 * 60 # 60 минут
775
+
776
+ for epoch in range(start_epoch, start_epoch + num_epochs):
777
+ batch_losses = []
778
+ batch_grads = []
779
+ batch_sampler.set_epoch(epoch)
780
+ accelerator.wait_for_everyone()
781
+ unet.train()
782
+
783
+ for step, (latents, embeddings, attention_mask, pooled) in enumerate(dataloader):
784
+ with accelerator.accumulate(unet):
785
+ if save_model == False and epoch == 0 and step == 5 :
786
+ used_gb = torch.cuda.max_memory_allocated() / 1024**3
787
+ print(f"Шаг {step}: {used_gb:.2f} GB")
788
+
789
+ # шум
790
+ noise = torch.randn_like(latents, dtype=latents.dtype)
791
+
792
+ # 3. Время t (сэмплим, как и раньше, но чуть сжимаем края)
793
+ u = torch.rand(latents.shape[0], device=latents.device, dtype=latents.dtype)
794
+ t = u * (1 - 2 * 1e-5) + 1e-5 # Теперь t строго в (0.00001 ... 0.99999)
795
+ # интерполяция между x0 и шумом
796
+ noisy_latents = (1.0 - t.view(-1, 1, 1, 1)) * latents + t.view(-1, 1, 1, 1) * noise
797
+ # делаем integer timesteps для UNet
798
+ timesteps = t.to(torch.float32).mul(999.0)
799
+ timesteps = timesteps.clamp(0, scheduler.config.num_train_timesteps - 1)
800
+
801
+ time_ids = torch.zeros(
802
+ pooled.shape[0], # ← вот это главное
803
+ 6,
804
+ device=device,
805
+ dtype=torch.long
806
+ )
807
+
808
+ # --- Вызов UNet с маской ---
809
+ model_pred = unet(
810
+ noisy_latents,
811
+ timesteps,
812
+ encoder_hidden_states=embeddings,
813
+ encoder_attention_mask=attention_mask,
814
+ added_cond_kwargs={"text_embeds": pooled,"time_ids": time_ids},
815
+ ).sample
816
+
817
+ target = noise - latents
818
+
819
+ mse_loss = F.mse_loss(model_pred.float(), target.float())
820
+ batch_losses.append(mse_loss.detach().item())
821
+
822
+ if (global_step % 100 == 0) or (global_step % sample_interval == 0):
823
+ accelerator.wait_for_everyone()
824
+
825
+ losses_dict = {}
826
+ losses_dict["mse"] = mse_loss
827
+
828
+ if (global_step % 100 == 0) or (global_step % sample_interval == 0):
829
+ accelerator.wait_for_everyone()
830
+
831
+ accelerator.backward(mse_loss)
832
+
833
+ if (global_step % 100 == 0) or (global_step % sample_interval == 0):
834
+ accelerator.wait_for_everyone()
835
+
836
+ grad = 0.0
837
+ if not fbp:
838
+ if accelerator.sync_gradients:
839
+ grad_val = accelerator.clip_grad_norm_(unet.parameters(), clip_grad_norm)
840
+ grad = grad_val.float().item() if torch.is_tensor(grad_val) else float(grad_val)
841
+ optimizer.step()
842
+ lr_scheduler.step()
843
+ optimizer.zero_grad(set_to_none=True)
844
+
845
+ if accelerator.sync_gradients:
846
+ global_step += 1
847
+ progress_bar.update(1)
848
+ if accelerator.is_main_process:
849
+ if fbp:
850
+ current_lr = base_learning_rate
851
+ else:
852
+ current_lr = lr_scheduler.get_last_lr()[0]
853
+ batch_grads.append(grad)
854
+
855
+ log_data = {}
856
+ log_data["loss_mse"] = mse_loss.detach().item()
857
+ log_data["lr"] = current_lr
858
+ log_data["grad"] = grad
859
+ if accelerator.sync_gradients:
860
+ if use_wandb:
861
+ wandb.log(log_data, step=global_step)
862
+ if use_comet_ml:
863
+ comet_experiment.log_metrics(log_data, step=global_step)
864
+
865
+ current_time = time.time()
866
+ is_time_to_sample = (current_time - last_sample_time) >= sample_interval_seconds
867
+ if is_time_to_sample or global_step == 50:
868
+ # Передаем tuple (emb, mask) для негатива
869
+ if save_model:
870
+ generate_and_save_samples(fixed_samples, (uncond_emb, uncond_mask,uncond_pooled), global_step)
871
+ elif epoch % 10 == 0:
872
+ generate_and_save_samples(fixed_samples, (uncond_emb, uncond_mask,uncond_pooled), global_step)
873
+ last_n = sample_interval
874
+
875
+ if save_model:
876
+ has_losses = len(batch_losses) > 0
877
+ avg_sample_loss = np.mean(batch_losses[-sample_interval:]) if has_losses else 0.0
878
+ last_loss = batch_losses[-1] if has_losses else 0.0
879
+ max_loss = max(avg_sample_loss, last_loss)
880
+ should_save = max_loss < min_loss * save_barrier
881
+ print(
882
+ f"Saving: {should_save} | Max: {max_loss:.4f} | "
883
+ f"Last: {last_loss:.4f} | Avg: {avg_sample_loss:.4f}"
884
+ )
885
+ # 6. Сохранение и обновление
886
+ if should_save:
887
+ min_loss = max_loss
888
+ save_checkpoint(unet)
889
+ last_sample_time = current_time
890
+ unet.train()
891
+
892
+ if accelerator.is_main_process:
893
+ avg_epoch_loss = np.mean(batch_losses) if len(batch_losses) > 0 else 0.0
894
+ avg_epoch_grad = np.mean(batch_grads) if len(batch_grads) > 0 else 0.0
895
+
896
+ print(f"\nЭпоха {epoch} завершена. Средний лосс: {avg_epoch_loss:.6f}")
897
+ log_data_ep = {
898
+ "epoch_loss": avg_epoch_loss,
899
+ "epoch_grad": avg_epoch_grad,
900
+ "epoch": epoch + 1,
901
+ }
902
+ if use_wandb:
903
+ wandb.log(log_data_ep)
904
+ if use_comet_ml:
905
+ comet_experiment.log_metrics(log_data_ep)
906
+
907
+ if accelerator.is_main_process:
908
+ print("Обучение завершено! Сохраняем финальную модель...")
909
+ #if save_model:
910
+ save_checkpoint(unet,"fp16")
911
+ if use_comet_ml:
912
+ comet_experiment.end()
913
+ accelerator.free_memory()
914
+ if torch.distributed.is_initialized():
915
+ torch.distributed.destroy_process_group()
916
+
917
+ print("Готово!")
train.py CHANGED
@@ -47,8 +47,8 @@ num_epochs = num_gpus #8 * max(1, int(num_gpus / 2))
47
  sample_interval_share = 20
48
  cfg_dropout = 0.10
49
  max_length = 248
50
- use_wandb = False
51
- use_comet_ml = True
52
  save_model = True
53
  use_decay = True
54
  fbp = False
@@ -193,8 +193,20 @@ def encode_texts(texts, max_length=max_length):
193
  last_hidden = outputs.hidden_states[layer_index]
194
  seq_len = toks.attention_mask.sum(dim=1) - 1
195
  pooled = last_hidden[torch.arange(len(last_hidden)), seq_len.clamp(min=0)]
196
- #pooled = torch.cat([pooled_clip, pooled], dim=1)
197
- return last_hidden.to(dtype), toks.attention_mask, pooled.to(dtype)
 
 
 
 
 
 
 
 
 
 
 
 
198
 
199
  shift_factor = getattr(vae.config, "shift_factor", 0.0)
200
  if shift_factor is None:
@@ -408,9 +420,9 @@ def get_fixed_samples_by_resolution(dataset, samples_per_group=1):
408
  texts = [item["text"] for item in samples_data]
409
 
410
  # Кодируем тексты на лету, чтобы получить маски и пулинг
411
- embeddings, masks, pooled = encode_texts(texts)
412
 
413
- fixed_samples[size] = (latents, embeddings, masks, texts, pooled)
414
 
415
  print(f"Создано {len(fixed_samples)} групп фиксированных семплов по разрешениям")
416
  return fixed_samples
@@ -446,12 +458,12 @@ def collate_fn_simple(batch):
446
  ]
447
  # 3. Кодируем на лету
448
  # Возвращает: hidden (B, L, D), mask (B, L)
449
- embeddings, attention_mask, pooled = encode_texts(texts)
450
 
451
  # attention_mask от токенизатора уже имеет нужный формат, но на всякий случай приведем к long
452
  attention_mask = attention_mask.to(dtype=torch.int64)
453
 
454
- return latents, embeddings, attention_mask, pooled
455
 
456
  batch_sampler = DistributedResolutionBatchSampler(
457
  dataset=dataset,
@@ -558,21 +570,20 @@ def get_negative_embedding(neg_prompt="", batch_size=1):
558
  empty_mask = torch.ones((batch_size, seq_len), dtype=torch.int64, device=device)
559
  return empty_emb, empty_mask
560
 
561
- uncond_emb, uncond_mask, uncond_pooled = encode_texts([neg_prompt])
562
  uncond_emb = uncond_emb.to(dtype=dtype, device=device).repeat(batch_size, 1, 1)
563
  uncond_mask = uncond_mask.to(device=device).repeat(batch_size, 1)
564
- uncond_pooled = uncond_pooled.to(device=device).repeat(batch_size, 1)
565
 
566
- return uncond_emb, uncond_mask, uncond_pooled
567
 
568
  # Получаем негативные (пустые) условия для валидации
569
- uncond_emb, uncond_mask, uncond_pooled = get_negative_embedding("low quality")
570
 
571
  # --- Функция генерации семплов ---
572
  @torch.compiler.disable()
573
  @torch.no_grad()
574
  def generate_and_save_samples(fixed_samples_cpu, uncond_data, step):
575
- uncond_emb, uncond_mask, uncond_pooled = uncond_data
576
 
577
  original_model = None
578
  try:
@@ -587,12 +598,11 @@ def generate_and_save_samples(fixed_samples_cpu, uncond_data, step):
587
  all_captions = []
588
 
589
  # Распаковываем 5 элементов (добавились mask)
590
- for size, (sample_latents, sample_text_embeddings, sample_mask, sample_text, sample_pooled) in fixed_samples_cpu.items():
591
  width, height = size
592
  sample_latents = sample_latents.to(dtype=dtype, device=device)
593
  sample_text_embeddings = sample_text_embeddings.to(dtype=dtype, device=device)
594
  sample_mask = sample_mask.to(device=device)
595
- sample_pooled = sample_pooled.to(dtype=dtype, device=device)
596
 
597
  latents = torch.randn(
598
  sample_latents.shape,
@@ -603,13 +613,6 @@ def generate_and_save_samples(fixed_samples_cpu, uncond_data, step):
603
 
604
  scheduler.set_timesteps(n_diffusion_steps, device=device)
605
 
606
- time_ids = torch.zeros(
607
- sample_pooled.shape[0], # ← вот это главное
608
- 6,
609
- device=device,
610
- dtype=torch.long
611
- )
612
-
613
  for t in scheduler.timesteps:
614
  if guidance_scale != 1:
615
  latent_model_input = torch.cat([latents, latents], dim=0)
@@ -624,29 +627,10 @@ def generate_and_save_samples(fixed_samples_cpu, uncond_data, step):
624
  neg_mask_batch = uncond_mask[0:1].expand(curr_batch_size, -1)
625
  attention_mask_batch = torch.cat([neg_mask_batch, sample_mask], dim=0)
626
 
627
- neg_pooled_batch = uncond_pooled[0:1].expand(curr_batch_size, -1)
628
- pooled_batch = torch.cat([neg_pooled_batch, sample_pooled], dim=0)
629
-
630
- # ← КЛЮЧЕВОЕ ИСПРАВЛЕНИЕ — time_ids под текущий удвоенный батч!
631
- time_ids = torch.zeros(
632
- pooled_batch.shape[0], # 2 * curr_batch_size при CFG
633
- 6,
634
- device=device,
635
- dtype=torch.long
636
- )
637
-
638
  else:
639
  latent_model_input = latents
640
  text_embeddings_batch = sample_text_embeddings
641
  attention_mask_batch = sample_mask
642
- pooled_batch = sample_pooled
643
-
644
- time_ids = torch.zeros(
645
- pooled_batch.shape[0],
646
- 6,
647
- device=device,
648
- dtype=torch.long
649
- )
650
 
651
  # Теперь всё имеет одинаковый batch size
652
  model_out = original_model(
@@ -654,10 +638,6 @@ def generate_and_save_samples(fixed_samples_cpu, uncond_data, step):
654
  t,
655
  encoder_hidden_states=text_embeddings_batch,
656
  encoder_attention_mask=attention_mask_batch,
657
- added_cond_kwargs={
658
- "text_embeds": pooled_batch,
659
- "time_ids": time_ids
660
- },
661
  )
662
 
663
  flow = getattr(model_out, "sample", model_out)
@@ -722,7 +702,7 @@ def generate_and_save_samples(fixed_samples_cpu, uncond_data, step):
722
  del all_generated_images, all_captions
723
  del latents, current_latents, latent_model_input, flow
724
  del decoded, decoded_fp32
725
- del sample_latents, sample_text_embeddings, sample_mask, sample_pooled # Копии на GPU
726
  del model_out
727
  except UnboundLocalError:
728
  pass
@@ -737,7 +717,7 @@ def generate_and_save_samples(fixed_samples_cpu, uncond_data, step):
737
  if accelerator.is_main_process:
738
  if save_model:
739
  print("Генерация сэмплов до старта обучения...")
740
- generate_and_save_samples(fixed_samples, (uncond_emb, uncond_mask,uncond_pooled), 0)
741
  accelerator.wait_for_everyone()
742
 
743
  def save_checkpoint(unet, variant=""):
@@ -780,7 +760,7 @@ for epoch in range(start_epoch, start_epoch + num_epochs):
780
  accelerator.wait_for_everyone()
781
  unet.train()
782
 
783
- for step, (latents, embeddings, attention_mask, pooled) in enumerate(dataloader):
784
  with accelerator.accumulate(unet):
785
  if save_model == False and epoch == 0 and step == 5 :
786
  used_gb = torch.cuda.max_memory_allocated() / 1024**3
@@ -798,20 +778,12 @@ for epoch in range(start_epoch, start_epoch + num_epochs):
798
  timesteps = t.to(torch.float32).mul(999.0)
799
  timesteps = timesteps.clamp(0, scheduler.config.num_train_timesteps - 1)
800
 
801
- time_ids = torch.zeros(
802
- pooled.shape[0], # ← вот это главное
803
- 6,
804
- device=device,
805
- dtype=torch.long
806
- )
807
-
808
  # --- Вызов UNet с маской ---
809
  model_pred = unet(
810
  noisy_latents,
811
  timesteps,
812
  encoder_hidden_states=embeddings,
813
  encoder_attention_mask=attention_mask,
814
- added_cond_kwargs={"text_embeds": pooled,"time_ids": time_ids},
815
  ).sample
816
 
817
  target = noise - latents
@@ -867,9 +839,9 @@ for epoch in range(start_epoch, start_epoch + num_epochs):
867
  if is_time_to_sample or global_step == 50:
868
  # Передаем tuple (emb, mask) для негатива
869
  if save_model:
870
- generate_and_save_samples(fixed_samples, (uncond_emb, uncond_mask,uncond_pooled), global_step)
871
  elif epoch % 10 == 0:
872
- generate_and_save_samples(fixed_samples, (uncond_emb, uncond_mask,uncond_pooled), global_step)
873
  last_n = sample_interval
874
 
875
  if save_model:
 
47
  sample_interval_share = 20
48
  cfg_dropout = 0.10
49
  max_length = 248
50
+ use_wandb = True
51
+ use_comet_ml = False
52
  save_model = True
53
  use_decay = True
54
  fbp = False
 
193
  last_hidden = outputs.hidden_states[layer_index]
194
  seq_len = toks.attention_mask.sum(dim=1) - 1
195
  pooled = last_hidden[torch.arange(len(last_hidden)), seq_len.clamp(min=0)]
196
+ # --- НОВАЯ ЛОГИКА: ОБЪЕДИНЕНИЕ ДЛЯ КРОСС-ВНИМАНИЯ ---
197
+ # 1. Расширяем пулинг-вектор до последовательности [B, 1, 1024]
198
+ pooled_expanded = pooled.unsqueeze(1)
199
+
200
+ # 2. Объединяем последовательность токенов и пулинг-вектор
201
+ # !!! ИЗМЕНЕНИЕ ЗДЕСЬ !!!: Пулинг идет ПЕРВЫМ
202
+ # Теперь: [B, 1 + L, 1024]. Пулинг стал токеном в НАЧАЛЕ.
203
+ new_encoder_hidden_states = torch.cat([pooled_expanded, last_hidden], dim=1)
204
+
205
+ # 3. Обновляем маску внимания для нового токена
206
+ # Маска внимания: [B, 1 + L]. Добавляем 1 в НАЧАЛО.
207
+ # torch.ones((batch_size, 1), device=device) создает маску [B, 1] со значениями 1.
208
+ new_attention_mask = torch.cat([torch.ones((last_hidden.shape[0], 1), device=device), toks.attention_mask], dim=1)
209
+ return new_encoder_hidden_states.to(dtype), new_attention_mask
210
 
211
  shift_factor = getattr(vae.config, "shift_factor", 0.0)
212
  if shift_factor is None:
 
420
  texts = [item["text"] for item in samples_data]
421
 
422
  # Кодируем тексты на лету, чтобы получить маски и пулинг
423
+ embeddings, masks = encode_texts(texts)
424
 
425
+ fixed_samples[size] = (latents, embeddings, masks, texts)
426
 
427
  print(f"Создано {len(fixed_samples)} групп фиксированных семплов по разрешениям")
428
  return fixed_samples
 
458
  ]
459
  # 3. Кодируем на лету
460
  # Возвращает: hidden (B, L, D), mask (B, L)
461
+ embeddings, attention_mask = encode_texts(texts)
462
 
463
  # attention_mask от токенизатора уже имеет нужный формат, но на всякий случай приведем к long
464
  attention_mask = attention_mask.to(dtype=torch.int64)
465
 
466
+ return latents, embeddings, attention_mask
467
 
468
  batch_sampler = DistributedResolutionBatchSampler(
469
  dataset=dataset,
 
570
  empty_mask = torch.ones((batch_size, seq_len), dtype=torch.int64, device=device)
571
  return empty_emb, empty_mask
572
 
573
+ uncond_emb, uncond_mask = encode_texts([neg_prompt])
574
  uncond_emb = uncond_emb.to(dtype=dtype, device=device).repeat(batch_size, 1, 1)
575
  uncond_mask = uncond_mask.to(device=device).repeat(batch_size, 1)
 
576
 
577
+ return uncond_emb, uncond_mask
578
 
579
  # Получаем негативные (пустые) условия для валидации
580
+ uncond_emb, uncond_mask = get_negative_embedding("low quality")
581
 
582
  # --- Функция генерации семплов ---
583
  @torch.compiler.disable()
584
  @torch.no_grad()
585
  def generate_and_save_samples(fixed_samples_cpu, uncond_data, step):
586
+ uncond_emb, uncond_mask = uncond_data
587
 
588
  original_model = None
589
  try:
 
598
  all_captions = []
599
 
600
  # Распаковываем 5 элементов (добавились mask)
601
+ for size, (sample_latents, sample_text_embeddings, sample_mask, sample_text) in fixed_samples_cpu.items():
602
  width, height = size
603
  sample_latents = sample_latents.to(dtype=dtype, device=device)
604
  sample_text_embeddings = sample_text_embeddings.to(dtype=dtype, device=device)
605
  sample_mask = sample_mask.to(device=device)
 
606
 
607
  latents = torch.randn(
608
  sample_latents.shape,
 
613
 
614
  scheduler.set_timesteps(n_diffusion_steps, device=device)
615
 
 
 
 
 
 
 
 
616
  for t in scheduler.timesteps:
617
  if guidance_scale != 1:
618
  latent_model_input = torch.cat([latents, latents], dim=0)
 
627
  neg_mask_batch = uncond_mask[0:1].expand(curr_batch_size, -1)
628
  attention_mask_batch = torch.cat([neg_mask_batch, sample_mask], dim=0)
629
 
 
 
 
 
 
 
 
 
 
 
 
630
  else:
631
  latent_model_input = latents
632
  text_embeddings_batch = sample_text_embeddings
633
  attention_mask_batch = sample_mask
 
 
 
 
 
 
 
 
634
 
635
  # Теперь всё имеет одинаковый batch size
636
  model_out = original_model(
 
638
  t,
639
  encoder_hidden_states=text_embeddings_batch,
640
  encoder_attention_mask=attention_mask_batch,
 
 
 
 
641
  )
642
 
643
  flow = getattr(model_out, "sample", model_out)
 
702
  del all_generated_images, all_captions
703
  del latents, current_latents, latent_model_input, flow
704
  del decoded, decoded_fp32
705
+ del sample_latents, sample_text_embeddings, sample_mask # Копии на GPU
706
  del model_out
707
  except UnboundLocalError:
708
  pass
 
717
  if accelerator.is_main_process:
718
  if save_model:
719
  print("Генерация сэмплов до старта обучения...")
720
+ generate_and_save_samples(fixed_samples, (uncond_emb, uncond_mask), 0)
721
  accelerator.wait_for_everyone()
722
 
723
  def save_checkpoint(unet, variant=""):
 
760
  accelerator.wait_for_everyone()
761
  unet.train()
762
 
763
+ for step, (latents, embeddings, attention_mask) in enumerate(dataloader):
764
  with accelerator.accumulate(unet):
765
  if save_model == False and epoch == 0 and step == 5 :
766
  used_gb = torch.cuda.max_memory_allocated() / 1024**3
 
778
  timesteps = t.to(torch.float32).mul(999.0)
779
  timesteps = timesteps.clamp(0, scheduler.config.num_train_timesteps - 1)
780
 
 
 
 
 
 
 
 
781
  # --- Вызов UNet с маской ---
782
  model_pred = unet(
783
  noisy_latents,
784
  timesteps,
785
  encoder_hidden_states=embeddings,
786
  encoder_attention_mask=attention_mask,
 
787
  ).sample
788
 
789
  target = noise - latents
 
839
  if is_time_to_sample or global_step == 50:
840
  # Передаем tuple (emb, mask) для негатива
841
  if save_model:
842
+ generate_and_save_samples(fixed_samples, (uncond_emb, uncond_mask), global_step)
843
  elif epoch % 10 == 0:
844
+ generate_and_save_samples(fixed_samples, (uncond_emb, uncond_mask), global_step)
845
  last_n = sample_interval
846
 
847
  if save_model:
unet/config.json CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:8c7b2a5355cb645ab748031633b38153e7c01ecef1cf512e781dfc9536e3892e
3
- size 1885
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0a85ea1867dbee11485b2de5f5777cf16f5c5a2ed261dba0a465f5c649092299
3
+ size 1879
unet/diffusion_pytorch_model.safetensors CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:de4e90a0bfd3a0448c78356a6e727f406e56f8517dbd2dc2d448a69b139fb9b3
3
- size 6318956752
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:aabe27ef505108f4dacbc87261a281295c784dca817224c8414ce87cee90a534
3
+ size 6294042336
unet1.5b-2TE-text-Copy1.ipynb ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:48c27ed007de2c0e26790ac0ef7ced4511b6aebeaa3b9e8abfa984f5acb62a3d
3
+ size 70150
unet1.5b-2TE-text.ipynb CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:48c27ed007de2c0e26790ac0ef7ced4511b6aebeaa3b9e8abfa984f5acb62a3d
3
- size 70150
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0e8e3028e9acfe5c8bf1cf2cb3a371eb91405c8080a120510172768bd86009ba
3
+ size 45191