Text-to-Image
Diffusers
Safetensors
recoilme commited on
Commit
445fd7f
·
1 Parent(s): 5b6a615
Files changed (47) hide show
  1. pipeline_sdxs-Copy1.py +176 -0
  2. pipeline_sdxs.py +62 -32
  3. samples/unet1b_320x640_0.jpg +3 -0
  4. samples/unet1b_352x640_0.jpg +3 -0
  5. samples/unet1b_384x640_0.jpg +3 -0
  6. samples/unet1b_416x640_0.jpg +3 -0
  7. samples/unet1b_448x640_0.jpg +3 -0
  8. samples/unet1b_480x640_0.jpg +3 -0
  9. samples/unet1b_512x640_0.jpg +3 -0
  10. samples/unet1b_544x640_0.jpg +3 -0
  11. samples/unet1b_576x640_0.jpg +3 -0
  12. samples/unet1b_608x640_0.jpg +3 -0
  13. samples/unet1b_640x320_0.jpg +3 -0
  14. samples/unet1b_640x352_0.jpg +3 -0
  15. samples/unet1b_640x384_0.jpg +3 -0
  16. samples/unet1b_640x416_0.jpg +3 -0
  17. samples/unet1b_640x448_0.jpg +3 -0
  18. samples/unet1b_640x480_0.jpg +3 -0
  19. samples/unet1b_640x512_0.jpg +3 -0
  20. samples/unet1b_640x544_0.jpg +3 -0
  21. samples/unet1b_640x576_0.jpg +3 -0
  22. samples/unet1b_640x608_0.jpg +3 -0
  23. samples/unet1b_640x640_0.jpg +3 -0
  24. samples/unet_320x640_0.jpg +2 -2
  25. samples/unet_352x640_0.jpg +2 -2
  26. samples/unet_384x640_0.jpg +2 -2
  27. samples/unet_416x640_0.jpg +2 -2
  28. samples/unet_448x640_0.jpg +2 -2
  29. samples/unet_480x640_0.jpg +2 -2
  30. samples/unet_512x640_0.jpg +2 -2
  31. samples/unet_544x640_0.jpg +2 -2
  32. samples/unet_576x640_0.jpg +2 -2
  33. samples/unet_608x640_0.jpg +2 -2
  34. samples/unet_640x320_0.jpg +2 -2
  35. samples/unet_640x352_0.jpg +2 -2
  36. samples/unet_640x384_0.jpg +2 -2
  37. samples/unet_640x416_0.jpg +2 -2
  38. samples/unet_640x448_0.jpg +2 -2
  39. samples/unet_640x480_0.jpg +2 -2
  40. samples/unet_640x512_0.jpg +2 -2
  41. samples/unet_640x544_0.jpg +2 -2
  42. samples/unet_640x576_0.jpg +2 -2
  43. samples/unet_640x608_0.jpg +2 -2
  44. samples/unet_640x640_0.jpg +2 -2
  45. train.py +19 -18
  46. unet/diffusion_pytorch_model.safetensors +2 -2
  47. unet1b.ipynb +3 -0
pipeline_sdxs-Copy1.py ADDED
@@ -0,0 +1,176 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import numpy as np
3
+ from PIL import Image
4
+ from typing import List, Union, Optional, Tuple
5
+ from dataclasses import dataclass
6
+
7
+ from diffusers import DiffusionPipeline
8
+ from diffusers.utils import BaseOutput
9
+ from tqdm import tqdm
10
+
11
+ @dataclass
12
+ class SdxsPipelineOutput(BaseOutput):
13
+ images: Union[List[Image.Image], np.ndarray]
14
+
15
+ class SdxsPipeline(DiffusionPipeline):
16
+ def __init__(self, vae, text_encoder, tokenizer, unet, scheduler):
17
+ super().__init__()
18
+ self.register_modules(
19
+ vae=vae,
20
+ text_encoder=text_encoder,
21
+ tokenizer=tokenizer,
22
+ unet=unet,
23
+ scheduler=scheduler
24
+ )
25
+ self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
26
+
27
+ def encode_prompt(self, prompt, negative_prompt, device, dtype):
28
+ """
29
+ Полное соответствие функции encode_texts и get_negative_embedding из трейна.
30
+ """
31
+ def get_single_encode(texts, is_negative=False):
32
+ if texts is None or texts == "":
33
+ # Логика get_negative_embedding из трейна
34
+ hidden_dim = self.text_encoder.config.hidden_size
35
+
36
+ shape = (1, self.text_encoder.config.max_position_embeddings, hidden_dim)
37
+ # В трейне для негатива: zeros для эмбеддингов и ones для маски
38
+ emb = torch.zeros(shape, dtype=dtype, device=device)
39
+ mask = torch.ones((1, self.text_encoder.config.max_position_embeddings), dtype=torch.int64, device=device)
40
+ return emb, mask
41
+
42
+ if isinstance(texts, str):
43
+ texts = [texts]
44
+
45
+ with torch.no_grad():
46
+ toks = self.tokenizer(
47
+ texts,
48
+ padding="max_length",
49
+ max_length=self.text_encoder.config.max_position_embeddings,
50
+ truncation=True,
51
+ return_tensors="pt"
52
+ ).to(device)
53
+
54
+ outputs = self.text_encoder(
55
+ input_ids=toks.input_ids,
56
+ attention_mask=toks.attention_mask,
57
+ output_hidden_states=True
58
+ )
59
+
60
+ # 1. Выбираем нужный слой.
61
+ # -1 — это последний блок трансформера
62
+ # -2 — это предпоследний (стандарт для большинства современных моделей)
63
+ layer_index = -2
64
+ prompt_embeds = outputs.hidden_states[layer_index]
65
+
66
+ # 2. ДОБАВЛЯЕМ ФИНАЛЬНУЮ НОРМАЛИЗАЦИЮ
67
+ # В CLIP после всех блоков стоит слой LayerNorm.
68
+ # Если мы берем скрытые состояния напрямую, мы "проскакиваем" его.
69
+ # Нужно применить его вручную:
70
+ final_layer_norm = self.text_encoder.text_model.final_layer_norm
71
+ prompt_embeds = final_layer_norm(prompt_embeds)
72
+
73
+ return prompt_embeds, toks.attention_mask
74
+
75
+ # Получаем эмбеддинги
76
+ pos_embeds, pos_mask = get_single_encode(prompt)
77
+ neg_embeds, neg_mask = get_single_encode(negative_prompt, is_negative=True)
78
+
79
+ # Выравнивание батча
80
+ batch_size = pos_embeds.shape[0]
81
+ if neg_embeds.shape[0] != batch_size:
82
+ neg_embeds = neg_embeds.repeat(batch_size, 1, 1)
83
+ neg_mask = neg_mask.repeat(batch_size, 1)
84
+
85
+ # Конкатенация для CFG: [Negative, Positive]
86
+ text_embeddings = torch.cat([neg_embeds, pos_embeds], dim=0)
87
+ final_mask = torch.cat([neg_mask, pos_mask], dim=0)
88
+
89
+ return text_embeddings.to(dtype=dtype), final_mask.to(dtype=torch.int64)
90
+
91
+ @torch.no_grad()
92
+ def __call__(
93
+ self,
94
+ prompt: Union[str, List[str]],
95
+ negative_prompt: Optional[Union[str, List[str]]] = None,
96
+ height: int = 1024,
97
+ width: int = 1024,
98
+ num_inference_steps: int = 40, # Как в трейне n_diffusion_steps
99
+ guidance_scale: float = 4.0, # Как в трейне
100
+ generator: Optional[torch.Generator] = None,
101
+ output_type: str = "pil",
102
+ return_dict: bool = True,
103
+ **kwargs,
104
+ ):
105
+ device = self.device
106
+ self.vae.to(device)
107
+ vae_scaling_factor = getattr(self.vae.config, "scaling_factor", 1.0)
108
+ vae_shift_factor = getattr(self.vae.config, "shift_factor", 0.0)
109
+
110
+ # 1. Encode Prompt
111
+ dtype = self.text_encoder.dtype
112
+ text_embeddings, attention_mask = self.encode_prompt(
113
+ prompt, negative_prompt, device, dtype
114
+ )
115
+
116
+ # 2. Prepare Latents
117
+ batch_size = 1 if isinstance(prompt, str) else len(prompt)
118
+ latent_channels = self.unet.config.in_channels
119
+
120
+ latents = torch.randn(
121
+ (batch_size, latent_channels, height // self.vae_scale_factor, width // self.vae_scale_factor),
122
+ generator=generator,
123
+ device=device,
124
+ dtype=dtype
125
+ )
126
+
127
+ # 3. Настройка Flow Matching шедулера
128
+ self.scheduler.set_timesteps(num_inference_steps, device=device)
129
+ timesteps = self.scheduler.timesteps
130
+
131
+ # 4. Denoising Loop
132
+ for t in tqdm(timesteps, desc="Sampling"):
133
+ # CFG input
134
+ latent_model_input = torch.cat([latents] * 2) if guidance_scale > 1 else latents
135
+
136
+ # Flow Matching обычно не требует scale_model_input,
137
+ # но оставим для совместимости с интерфейсом шедулера
138
+ if hasattr(self.scheduler, "scale_model_input"):
139
+ latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
140
+
141
+ # Predict
142
+ model_out = self.unet(
143
+ latent_model_input,
144
+ t,
145
+ encoder_hidden_states=text_embeddings,
146
+ encoder_attention_mask=attention_mask,
147
+ return_dict=False,
148
+ )[0]
149
+
150
+ # CFG Logic
151
+ if guidance_scale > 1:
152
+ flow_uncond, flow_cond = model_out.chunk(2)
153
+ model_out = flow_uncond + guidance_scale * (flow_cond - flow_uncond)
154
+
155
+ # Step (Flow Matching Euler)
156
+ latents = self.scheduler.step(model_out, t, latents, return_dict=False)[0]
157
+
158
+ # 5. Decode
159
+ if output_type == "latent":
160
+ return SdxsPipelineOutput(images=latents)
161
+
162
+ latents = latents * vae_scaling_factor + vae_shift_factor
163
+ image = self.vae.decode(latents.to(self.vae.dtype), return_dict=False)[0]
164
+
165
+ # Пост-процессинг
166
+ image = (image / 2 + 0.5).clamp(0, 1)
167
+ image = image.cpu().permute(0, 2, 3, 1).float().numpy()
168
+
169
+ if output_type == "pil":
170
+ image = (image * 255).round().astype("uint8")
171
+ image = [Image.fromarray(img) for img in image]
172
+
173
+ if not return_dict:
174
+ return image
175
+
176
+ return SdxsPipelineOutput(images=image)
pipeline_sdxs.py CHANGED
@@ -24,17 +24,31 @@ class SdxsPipeline(DiffusionPipeline):
24
  )
25
  self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
26
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
27
  def encode_prompt(self, prompt, negative_prompt, device, dtype):
28
- """
29
- Полное соответствие функции encode_texts и get_negative_embedding из трейна.
30
- """
31
  def get_single_encode(texts, is_negative=False):
32
  if texts is None or texts == "":
33
- # Логика get_negative_embedding из трейна
34
  hidden_dim = self.text_encoder.config.hidden_size
35
-
36
  shape = (1, self.text_encoder.config.max_position_embeddings, hidden_dim)
37
- # В трейне для негатива: zeros для эмбеддингов и ones для маски
38
  emb = torch.zeros(shape, dtype=dtype, device=device)
39
  mask = torch.ones((1, self.text_encoder.config.max_position_embeddings), dtype=torch.int64, device=device)
40
  return emb, mask
@@ -57,32 +71,21 @@ class SdxsPipeline(DiffusionPipeline):
57
  output_hidden_states=True
58
  )
59
 
60
- # 1. Выбираем нужный слой.
61
- # -1 — это последний блок трансформера
62
- # -2 — это предпоследний (стандарт для большинства современных моделей)
63
  layer_index = -2
64
  prompt_embeds = outputs.hidden_states[layer_index]
65
-
66
- # 2. ДОБАВЛЯЕМ ФИНАЛЬНУЮ НОРМАЛИЗАЦИЮ
67
- # В CLIP после всех блоков стоит слой LayerNorm.
68
- # Если мы берем скрытые состояния напрямую, мы "проскакиваем" его.
69
- # Нужно применить его вручную:
70
  final_layer_norm = self.text_encoder.text_model.final_layer_norm
71
  prompt_embeds = final_layer_norm(prompt_embeds)
72
 
73
  return prompt_embeds, toks.attention_mask
74
 
75
- # Получаем эмбеддинги
76
  pos_embeds, pos_mask = get_single_encode(prompt)
77
  neg_embeds, neg_mask = get_single_encode(negative_prompt, is_negative=True)
78
 
79
- # Выравнивание батча
80
  batch_size = pos_embeds.shape[0]
81
  if neg_embeds.shape[0] != batch_size:
82
  neg_embeds = neg_embeds.repeat(batch_size, 1, 1)
83
  neg_mask = neg_mask.repeat(batch_size, 1)
84
 
85
- # Конкатенация для CFG: [Negative, Positive]
86
  text_embeddings = torch.cat([neg_embeds, pos_embeds], dim=0)
87
  final_mask = torch.cat([neg_mask, pos_mask], dim=0)
88
 
@@ -92,11 +95,13 @@ class SdxsPipeline(DiffusionPipeline):
92
  def __call__(
93
  self,
94
  prompt: Union[str, List[str]],
 
 
95
  negative_prompt: Optional[Union[str, List[str]]] = None,
96
  height: int = 1024,
97
  width: int = 1024,
98
- num_inference_steps: int = 40, # Как в трейне n_diffusion_steps
99
- guidance_scale: float = 4.0, # Как в трейне
100
  generator: Optional[torch.Generator] = None,
101
  output_type: str = "pil",
102
  return_dict: bool = True,
@@ -113,28 +118,53 @@ class SdxsPipeline(DiffusionPipeline):
113
  prompt, negative_prompt, device, dtype
114
  )
115
 
116
- # 2. Prepare Latents
117
  batch_size = 1 if isinstance(prompt, str) else len(prompt)
118
  latent_channels = self.unet.config.in_channels
119
-
120
- latents = torch.randn(
121
- (batch_size, latent_channels, height // self.vae_scale_factor, width // self.vae_scale_factor),
122
- generator=generator,
123
- device=device,
124
- dtype=dtype
125
- )
126
 
127
- # 3. Настройка Flow Matching шедулера
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
128
  self.scheduler.set_timesteps(num_inference_steps, device=device)
129
  timesteps = self.scheduler.timesteps
130
 
 
 
 
 
 
 
131
  # 4. Denoising Loop
132
  for t in tqdm(timesteps, desc="Sampling"):
133
- # CFG input
134
  latent_model_input = torch.cat([latents] * 2) if guidance_scale > 1 else latents
135
 
136
- # Flow Matching обычно не требует scale_model_input,
137
- # но оставим для совместимости с интерфейсом шедулера
138
  if hasattr(self.scheduler, "scale_model_input"):
139
  latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
140
 
@@ -152,7 +182,7 @@ class SdxsPipeline(DiffusionPipeline):
152
  flow_uncond, flow_cond = model_out.chunk(2)
153
  model_out = flow_uncond + guidance_scale * (flow_cond - flow_uncond)
154
 
155
- # Step (Flow Matching Euler)
156
  latents = self.scheduler.step(model_out, t, latents, return_dict=False)[0]
157
 
158
  # 5. Decode
 
24
  )
25
  self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
26
 
27
+ # --- ВСПОМОГАТЕЛЬНАЯ ФУНКЦИЯ ДЛЯ ПОДГОТОВКИ ИЗОБРАЖЕНИЯ (Img2Img) ---
28
+ def preprocess_image(self, image: Image.Image, width: int, height: int):
29
+ """Ресайз и центрированный кроп изображения под нужный размер"""
30
+ w, h = image.size
31
+ aspect_ratio = width / height
32
+ if w / h > aspect_ratio:
33
+ new_w = int(h * aspect_ratio)
34
+ left = (w - new_w) // 2
35
+ image = image.crop((left, 0, left + new_w, h))
36
+ else:
37
+ new_h = int(w / aspect_ratio)
38
+ top = (h - new_h) // 2
39
+ image = image.crop((0, top, w, top + new_h))
40
+
41
+ image = image.resize((width, height), resample=Image.LANCZOS)
42
+ image = np.array(image).astype(np.float32) / 255.0
43
+ image = image[None].transpose(0, 3, 1, 2) # [1, C, H, W]
44
+ image = torch.from_numpy(image)
45
+ return 2.0 * image - 1.0 # В диапазон [-1, 1]
46
+
47
  def encode_prompt(self, prompt, negative_prompt, device, dtype):
 
 
 
48
  def get_single_encode(texts, is_negative=False):
49
  if texts is None or texts == "":
 
50
  hidden_dim = self.text_encoder.config.hidden_size
 
51
  shape = (1, self.text_encoder.config.max_position_embeddings, hidden_dim)
 
52
  emb = torch.zeros(shape, dtype=dtype, device=device)
53
  mask = torch.ones((1, self.text_encoder.config.max_position_embeddings), dtype=torch.int64, device=device)
54
  return emb, mask
 
71
  output_hidden_states=True
72
  )
73
 
 
 
 
74
  layer_index = -2
75
  prompt_embeds = outputs.hidden_states[layer_index]
 
 
 
 
 
76
  final_layer_norm = self.text_encoder.text_model.final_layer_norm
77
  prompt_embeds = final_layer_norm(prompt_embeds)
78
 
79
  return prompt_embeds, toks.attention_mask
80
 
 
81
  pos_embeds, pos_mask = get_single_encode(prompt)
82
  neg_embeds, neg_mask = get_single_encode(negative_prompt, is_negative=True)
83
 
 
84
  batch_size = pos_embeds.shape[0]
85
  if neg_embeds.shape[0] != batch_size:
86
  neg_embeds = neg_embeds.repeat(batch_size, 1, 1)
87
  neg_mask = neg_mask.repeat(batch_size, 1)
88
 
 
89
  text_embeddings = torch.cat([neg_embeds, pos_embeds], dim=0)
90
  final_mask = torch.cat([neg_mask, pos_mask], dim=0)
91
 
 
95
  def __call__(
96
  self,
97
  prompt: Union[str, List[str]],
98
+ image: Optional[Union[Image.Image, List[Image.Image]]] = None, # Добавлен параметр изображения
99
+ coef: float = 0.5, # Коэффициент влияния (strength): 1.0 - полный шум, 0.0 - оригинал
100
  negative_prompt: Optional[Union[str, List[str]]] = None,
101
  height: int = 1024,
102
  width: int = 1024,
103
+ num_inference_steps: int = 40,
104
+ guidance_scale: float = 4.0,
105
  generator: Optional[torch.Generator] = None,
106
  output_type: str = "pil",
107
  return_dict: bool = True,
 
118
  prompt, negative_prompt, device, dtype
119
  )
120
 
121
+ # 2. Определение размеров и батча
122
  batch_size = 1 if isinstance(prompt, str) else len(prompt)
123
  latent_channels = self.unet.config.in_channels
 
 
 
 
 
 
 
124
 
125
+ # --- ЛОГИКА ПОДГОТОВКИ ЛАТЕНТОВ (TXT2IMG vs IMG2IMG) ---
126
+ if image is not None:
127
+ # Если передана одна картинка на батч текстов - размножаем её
128
+ if isinstance(image, Image.Image):
129
+ image = [image] * batch_size
130
+
131
+ # Обработка картинок: ресайз/кроп и перевод в тензоры
132
+ image_tensors = torch.cat([self.preprocess_image(img, width, height) for img in image]).to(device=device, dtype=vae_dtype := self.vae.dtype)
133
+
134
+ # Кодируем в латенты
135
+ init_latents = self.vae.encode(image_tensors).latent_dist.sample(generator=generator)
136
+ init_latents = (init_latents - vae_shift_factor) / vae_scaling_factor
137
+ init_latents = init_latents.to(dtype=dtype)
138
+
139
+ # Создаем шум
140
+ noise = torch.randn(init_latents.shape, generator=generator, device=device, dtype=dtype)
141
+
142
+ # Flow Matching шум: x_t = (1 - t)*x0 + t*epsilon
143
+ # Здесь t = coef. Если coef=1.0, получаем чистый шум.
144
+ latents = (1.0 - coef) * init_latents + coef * noise
145
+ else:
146
+ # Стандартный txt2img (чистый шум)
147
+ latents = torch.randn(
148
+ (batch_size, latent_channels, height // self.vae_scale_factor, width // self.vae_scale_factor),
149
+ generator=generator,
150
+ device=device,
151
+ dtype=dtype
152
+ )
153
+
154
+ # 3. Настройка таймстепов
155
  self.scheduler.set_timesteps(num_inference_steps, device=device)
156
  timesteps = self.scheduler.timesteps
157
 
158
+ # --- ОБРЕЗКА ТАЙМСТЕПОВ ДЛЯ IMG2IMG ---
159
+ if image is not None:
160
+ # Оставляем только те шаги, которые соответствуют уровню шума coef
161
+ # В Flow Matching t идет от 1.0 к 0.0. Нам нужны шаги <= coef.
162
+ timesteps = timesteps[timesteps <= coef]
163
+
164
  # 4. Denoising Loop
165
  for t in tqdm(timesteps, desc="Sampling"):
 
166
  latent_model_input = torch.cat([latents] * 2) if guidance_scale > 1 else latents
167
 
 
 
168
  if hasattr(self.scheduler, "scale_model_input"):
169
  latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
170
 
 
182
  flow_uncond, flow_cond = model_out.chunk(2)
183
  model_out = flow_uncond + guidance_scale * (flow_cond - flow_uncond)
184
 
185
+ # Step
186
  latents = self.scheduler.step(model_out, t, latents, return_dict=False)[0]
187
 
188
  # 5. Decode
samples/unet1b_320x640_0.jpg ADDED

Git LFS Details

  • SHA256: d2aa4e330469a2c665ca049469544dca2539f4a96ec01a087f512c687635cad1
  • Pointer size: 131 Bytes
  • Size of remote file: 251 kB
samples/unet1b_352x640_0.jpg ADDED

Git LFS Details

  • SHA256: 5ca689118262aa05cae03886258c8dd27ef5694c3546207b8fabed21d866dd95
  • Pointer size: 131 Bytes
  • Size of remote file: 298 kB
samples/unet1b_384x640_0.jpg ADDED

Git LFS Details

  • SHA256: 618a920f14dbd164a24973364a053a5da229787dfa66169162bc32ad15f04402
  • Pointer size: 131 Bytes
  • Size of remote file: 389 kB
samples/unet1b_416x640_0.jpg ADDED

Git LFS Details

  • SHA256: 2657a51101a428fd44809a1d795156ddbfd593ad60f91db685a2b51f76efaa42
  • Pointer size: 131 Bytes
  • Size of remote file: 602 kB
samples/unet1b_448x640_0.jpg ADDED

Git LFS Details

  • SHA256: c2a13c181e0243849b6ca66b1b01db76f44ed27c84bbfb7e061c1bda41011ea0
  • Pointer size: 131 Bytes
  • Size of remote file: 310 kB
samples/unet1b_480x640_0.jpg ADDED

Git LFS Details

  • SHA256: 237b0891490d9645a70c5b7cb4d8827b78b1cb9ef5fb1dd5a461ba7bbae63a5e
  • Pointer size: 131 Bytes
  • Size of remote file: 267 kB
samples/unet1b_512x640_0.jpg ADDED

Git LFS Details

  • SHA256: 795dcfcbfdee4adde1dcb921a19da1741f7b136d76d6eae044e8419be50a053e
  • Pointer size: 131 Bytes
  • Size of remote file: 389 kB
samples/unet1b_544x640_0.jpg ADDED

Git LFS Details

  • SHA256: 5e4afd644ef6383745c0e254d07facaf5164ae555e3e853224bdbbcb6b5dccc8
  • Pointer size: 131 Bytes
  • Size of remote file: 611 kB
samples/unet1b_576x640_0.jpg ADDED

Git LFS Details

  • SHA256: c0a7c98f39d20e418a75fdb39f81b39bfbf89def0a50a42bf85a74ee9262c641
  • Pointer size: 131 Bytes
  • Size of remote file: 375 kB
samples/unet1b_608x640_0.jpg ADDED

Git LFS Details

  • SHA256: 5aff5dcbc205c957aedf5fce78dbc07d9bf34a7bb65d9ceb9140d420db8b0592
  • Pointer size: 131 Bytes
  • Size of remote file: 527 kB
samples/unet1b_640x320_0.jpg ADDED

Git LFS Details

  • SHA256: b9c52b10b0bb75838ed8449f0cf4b14b6f219248461c4b4f2f3b89f84b5ea8ce
  • Pointer size: 131 Bytes
  • Size of remote file: 168 kB
samples/unet1b_640x352_0.jpg ADDED

Git LFS Details

  • SHA256: 86a38d0237d797605208c15724c019f680c9484b27aded6da4efbdb7b9f3d874
  • Pointer size: 131 Bytes
  • Size of remote file: 243 kB
samples/unet1b_640x384_0.jpg ADDED

Git LFS Details

  • SHA256: b9a1f5bcb8faba50817680025dffe1013b0117e84a55ca92575db8bcf65beafb
  • Pointer size: 131 Bytes
  • Size of remote file: 320 kB
samples/unet1b_640x416_0.jpg ADDED

Git LFS Details

  • SHA256: 2e9856cd7c3e713218f33a14a7e2c315e8f76b655165887454b9748d3b8b6b45
  • Pointer size: 131 Bytes
  • Size of remote file: 479 kB
samples/unet1b_640x448_0.jpg ADDED

Git LFS Details

  • SHA256: 6cededbc1dee2c97e3e22c146bac13af9134e4369d84a2a1f89c5ce8dc09a4a9
  • Pointer size: 131 Bytes
  • Size of remote file: 448 kB
samples/unet1b_640x480_0.jpg ADDED

Git LFS Details

  • SHA256: 20bec749da0fda48aa72a909366f5ad9544b8f0e7927b4ac07dba0f733b650df
  • Pointer size: 131 Bytes
  • Size of remote file: 380 kB
samples/unet1b_640x512_0.jpg ADDED

Git LFS Details

  • SHA256: e0820ea4d654ad282f2990feae196efdc5a1f5807d46bbb1a77f44aa038c6d04
  • Pointer size: 131 Bytes
  • Size of remote file: 392 kB
samples/unet1b_640x544_0.jpg ADDED

Git LFS Details

  • SHA256: b1dc675eb4454d03a9aa4e058a50583ea37c5113e9e128f71634b8a87ac791d6
  • Pointer size: 131 Bytes
  • Size of remote file: 375 kB
samples/unet1b_640x576_0.jpg ADDED

Git LFS Details

  • SHA256: 6dba2421dacd33efad9d7efe8bfbdda46ae212903262f56c7a22cdf687d90e38
  • Pointer size: 131 Bytes
  • Size of remote file: 369 kB
samples/unet1b_640x608_0.jpg ADDED

Git LFS Details

  • SHA256: 17a304b89587795efdbf9e82dba47012297677660c9b1e4fab7109e8cde84a6e
  • Pointer size: 131 Bytes
  • Size of remote file: 853 kB
samples/unet1b_640x640_0.jpg ADDED

Git LFS Details

  • SHA256: 91b6dfa7c982eac7d91942abcf52883e1462ca342f4bee6f855021b0aa11985c
  • Pointer size: 131 Bytes
  • Size of remote file: 556 kB
samples/unet_320x640_0.jpg CHANGED

Git LFS Details

  • SHA256: e22bade07804d758475dfddc974c1568b9805cc2e7c30442f90f12238ae3af15
  • Pointer size: 131 Bytes
  • Size of remote file: 127 kB

Git LFS Details

  • SHA256: 860760bc841d0440d81221bff7f50ab55a622bc1a845709f1c1c7aac66010541
  • Pointer size: 131 Bytes
  • Size of remote file: 255 kB
samples/unet_352x640_0.jpg CHANGED

Git LFS Details

  • SHA256: a2d35433272a9fad677ba70bc4973977691a0310d6c07c1b7b4d2a5b5068cbe9
  • Pointer size: 131 Bytes
  • Size of remote file: 372 kB

Git LFS Details

  • SHA256: e9e0552ff1771aa832e839024726f51bc233d3df46e9f8193305f312cde8952c
  • Pointer size: 131 Bytes
  • Size of remote file: 174 kB
samples/unet_384x640_0.jpg CHANGED

Git LFS Details

  • SHA256: d5ddd73392ada85f3df31e6aade092a9ae0d3276f3652758f1d85f2c53421806
  • Pointer size: 131 Bytes
  • Size of remote file: 374 kB

Git LFS Details

  • SHA256: 7c0348cbab8a40013b960ebb4eaaa867b57d8ce2b7d5a26f8e09de9b5734209d
  • Pointer size: 131 Bytes
  • Size of remote file: 377 kB
samples/unet_416x640_0.jpg CHANGED

Git LFS Details

  • SHA256: 77d3b20229eafd6998ec9b544f3f8047303d6a249d80739f1eeca73633cc87c9
  • Pointer size: 131 Bytes
  • Size of remote file: 154 kB

Git LFS Details

  • SHA256: 1ed1020a57bd2bc244ca7f075c514235f3f6b4dae762608643a0c1cbcc055413
  • Pointer size: 131 Bytes
  • Size of remote file: 371 kB
samples/unet_448x640_0.jpg CHANGED

Git LFS Details

  • SHA256: cb1455c64024ab78d7a6e469568738b3bb0969626ae7fd5b550d4e08c18af478
  • Pointer size: 131 Bytes
  • Size of remote file: 316 kB

Git LFS Details

  • SHA256: bdd6a119d45bd8814f080dad7ec6550d35e0aad6ba5396b6d7a64eabc2723908
  • Pointer size: 131 Bytes
  • Size of remote file: 318 kB
samples/unet_480x640_0.jpg CHANGED

Git LFS Details

  • SHA256: 1f4312960dd18de33d6392ba11cf6a4e5be221b67e074955d9225d09a50883c5
  • Pointer size: 131 Bytes
  • Size of remote file: 436 kB

Git LFS Details

  • SHA256: 61e09d197abfd120212cc8594f83531fe4435ef9845196ff29ea9a021cff8c60
  • Pointer size: 131 Bytes
  • Size of remote file: 121 kB
samples/unet_512x640_0.jpg CHANGED

Git LFS Details

  • SHA256: 97b31083d874e903de42aa30f9cac87e42d750056d3f5bd4d8adce5066d03d31
  • Pointer size: 131 Bytes
  • Size of remote file: 344 kB

Git LFS Details

  • SHA256: f1e30fb64b9f91785e0cea6b1510e162adae122382b8faf3dbf22bc6436dda92
  • Pointer size: 131 Bytes
  • Size of remote file: 422 kB
samples/unet_544x640_0.jpg CHANGED

Git LFS Details

  • SHA256: 79e27f3a01187c3173c5e9dbd0b098b3e83f01d9ca19e1e07924360e6fa18c64
  • Pointer size: 132 Bytes
  • Size of remote file: 1.01 MB

Git LFS Details

  • SHA256: e1c27f5f0f7d5e83d0b721f9148c83a836b287d136fd096c293c465a602376ce
  • Pointer size: 131 Bytes
  • Size of remote file: 126 kB
samples/unet_576x640_0.jpg CHANGED

Git LFS Details

  • SHA256: f2c8e792ec50bba48ebf99854494d68c2574d86dc08f57153d3b67088a982810
  • Pointer size: 131 Bytes
  • Size of remote file: 196 kB

Git LFS Details

  • SHA256: 593c890911b74f33a234bcd9c828add2ea1d980c5ca285eea1422ea412b79530
  • Pointer size: 131 Bytes
  • Size of remote file: 228 kB
samples/unet_608x640_0.jpg CHANGED

Git LFS Details

  • SHA256: 3f9aae7894a780178d9a03f62ae169f0ac557f2742ea179312b0358ccdc388f7
  • Pointer size: 131 Bytes
  • Size of remote file: 289 kB

Git LFS Details

  • SHA256: 8f466cc2aaa8723cc45719bfa01ffbb7379b7ec8848f901f25de1142475d1ef2
  • Pointer size: 131 Bytes
  • Size of remote file: 402 kB
samples/unet_640x320_0.jpg CHANGED

Git LFS Details

  • SHA256: b3b9b6515b749df845735953981b6193c2559725a09ccddb965935cea52fa5e4
  • Pointer size: 131 Bytes
  • Size of remote file: 230 kB

Git LFS Details

  • SHA256: 644cc8ee0f6877396470a792dd5f00c2310cad50f25b0507e4d78748e9b20af2
  • Pointer size: 131 Bytes
  • Size of remote file: 261 kB
samples/unet_640x352_0.jpg CHANGED

Git LFS Details

  • SHA256: 63caf5249488fc3b4f114c4f3a396c942624ca6fb0642e178863db95a47c72c3
  • Pointer size: 131 Bytes
  • Size of remote file: 217 kB

Git LFS Details

  • SHA256: aedfeb4e3c265581c2665156c8bc7822b612110e0373e0de68efd0f46fcbc251
  • Pointer size: 131 Bytes
  • Size of remote file: 310 kB
samples/unet_640x384_0.jpg CHANGED

Git LFS Details

  • SHA256: 04fee8aa78a2b9c57dd359a018e27f08c7715755b6ae5e60e831776cd98e06cd
  • Pointer size: 131 Bytes
  • Size of remote file: 330 kB

Git LFS Details

  • SHA256: 9b6fe0726c1ad10aabe591bf3202d2fac094f2796c9073f90c195aa46343ccc3
  • Pointer size: 131 Bytes
  • Size of remote file: 206 kB
samples/unet_640x416_0.jpg CHANGED

Git LFS Details

  • SHA256: 247628793001819dada0065ccc237136ad55a110913da291ff17c06b5d162903
  • Pointer size: 131 Bytes
  • Size of remote file: 526 kB

Git LFS Details

  • SHA256: d851d5e1af3fad6b8f58a71b53cea550eaaf411bf5ece0ebb10056ad6a5f3e58
  • Pointer size: 131 Bytes
  • Size of remote file: 395 kB
samples/unet_640x448_0.jpg CHANGED

Git LFS Details

  • SHA256: 739debe89e7439a11b208b35ebfa42d419d7a745e38c60ca7478d32d0d26e414
  • Pointer size: 131 Bytes
  • Size of remote file: 250 kB

Git LFS Details

  • SHA256: 4f48e5f3f049b5313f866accff3066a6a42022a1b7db339df6a678e48176d7da
  • Pointer size: 131 Bytes
  • Size of remote file: 287 kB
samples/unet_640x480_0.jpg CHANGED

Git LFS Details

  • SHA256: 8a568b4181a9793d934e1bc620d879b126cf321103be0392719bd0347fc7fd65
  • Pointer size: 131 Bytes
  • Size of remote file: 336 kB

Git LFS Details

  • SHA256: ed29b73fe0ec5ae7a83b66e141f3a130cf416d0c790b91ff1aac4fe78f70206e
  • Pointer size: 131 Bytes
  • Size of remote file: 105 kB
samples/unet_640x512_0.jpg CHANGED

Git LFS Details

  • SHA256: bc80a1771b7004a74338b0e98ef0f5ad91588390c4a5c4729b6e4d89a4a44484
  • Pointer size: 130 Bytes
  • Size of remote file: 91.1 kB

Git LFS Details

  • SHA256: 8e8050747a9161f207d1e21ba0f659f58f9579a64074b1095fdf005c7b91f6ef
  • Pointer size: 131 Bytes
  • Size of remote file: 456 kB
samples/unet_640x544_0.jpg CHANGED

Git LFS Details

  • SHA256: aab84c1b9caeda7a66d0cab12639e2aff1177ec309da218251bbeaf069ded9c3
  • Pointer size: 131 Bytes
  • Size of remote file: 490 kB

Git LFS Details

  • SHA256: 2abfbae9a6a82c1e72a09ce58261b6ca146efe619cf3b9d49d49730fb162d40e
  • Pointer size: 131 Bytes
  • Size of remote file: 385 kB
samples/unet_640x576_0.jpg CHANGED

Git LFS Details

  • SHA256: 657a1aa18f51b1e18df8281446697a1f3c8735656112d07d78497cfe7b1dd1e8
  • Pointer size: 131 Bytes
  • Size of remote file: 889 kB

Git LFS Details

  • SHA256: b4d956147a40a7754e3a098043192a789cc76f20237af6b2f01e11b4f0d71881
  • Pointer size: 131 Bytes
  • Size of remote file: 767 kB
samples/unet_640x608_0.jpg CHANGED

Git LFS Details

  • SHA256: 6cabf716c74ec478eaf692fdd5647d075d45ee7b8a8e3a6b1f0062829fdb69dc
  • Pointer size: 132 Bytes
  • Size of remote file: 1.02 MB

Git LFS Details

  • SHA256: 3cd89416f52bbd90c81bc10c0b3ba713bf7cdf8da0d7b74fe46229faf76f6a5a
  • Pointer size: 132 Bytes
  • Size of remote file: 1.25 MB
samples/unet_640x640_0.jpg CHANGED

Git LFS Details

  • SHA256: e5f334e4551cb611405341dea74b8d7e207a32ae4a70edc874cf937cbecb3c38
  • Pointer size: 131 Bytes
  • Size of remote file: 615 kB

Git LFS Details

  • SHA256: b7ad0ea15f6f3912172769fc3d70219dfdb739baa7ab303aff51925b1e0ce877
  • Pointer size: 131 Bytes
  • Size of remote file: 681 kB
train.py CHANGED
@@ -1,5 +1,8 @@
1
  #from comet_ml import Experiment
2
  import os
 
 
 
3
  import math
4
  import torch
5
  import numpy as np
@@ -9,7 +12,7 @@ from torch.utils.data.distributed import DistributedSampler
9
  from torch.optim.lr_scheduler import LambdaLR
10
  from collections import defaultdict
11
  from diffusers import UNet2DConditionModel, AutoencoderKL,AutoencoderKLFlux2,AsymmetricAutoencoderKL,FlowMatchEulerDiscreteScheduler
12
- from accelerate import Accelerator
13
  from datasets import load_from_disk
14
  from tqdm import tqdm
15
  from PIL import Image, ImageOps
@@ -29,11 +32,12 @@ from transformers import AutoTokenizer, AutoModel
29
  # --------------------------- Параметры ---------------------------
30
  ds_path = "/workspace/sdxs-08b/datasets/d123_640_sd15"
31
  project = "unet"
32
- batch_size = 20 ## total batch (split // num `GPU)
 
33
  base_learning_rate = 2e-6
34
- min_learning_rate = 9e-7
35
- num_epochs = 3
36
- sample_interval_share = 5
37
  cfg_dropout = 0.10
38
  max_length = 248
39
  use_wandb = True
@@ -55,7 +59,6 @@ torch.backends.cudnn.allow_tf32 = True
55
  torch.backends.cuda.enable_flash_sdp(True)
56
  torch.backends.cuda.enable_mem_efficient_sdp(True)
57
  torch.backends.cuda.enable_math_sdp(False) # Отключаем медленный вариант
58
- dtype = torch.float32
59
  save_barrier = 1.05
60
  warmup_percent = 0.03
61
  percentile_clipping = 95
@@ -64,17 +67,9 @@ eps = 1e-7
64
  clip_grad_norm = 1.0
65
  limit = 0
66
  checkpoints_folder = ""
67
- mixed_precision = "no"
68
  gradient_accumulation_steps = 1
69
-
70
- # Это критично для новых карт, чтобы избежать фрагментации при больших контекстах
71
- os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True,max_split_size_mb:128"
72
-
73
- accelerator = Accelerator(
74
- mixed_precision=mixed_precision,
75
- gradient_accumulation_steps=gradient_accumulation_steps
76
- )
77
- device = accelerator.device
78
 
79
  # Параметры для диффузии
80
  n_diffusion_steps = 40
@@ -95,6 +90,12 @@ if fixed_seed:
95
  if torch.cuda.is_available():
96
  torch.cuda.manual_seed_all(seed)
97
 
 
 
 
 
 
 
98
  print("init")
99
 
100
  # --------------------------- Инициализация WandB ---------------------------
@@ -128,7 +129,7 @@ tokenizer = AutoTokenizer.from_pretrained("tokenizer")
128
  text_model = AutoModel.from_pretrained("text_encoder").to(device).eval()
129
  scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained("scheduler")
130
 
131
- def encode_texts(texts, max_length=max_length): # Для SD 1.5 лучше жестко 77
132
  if texts is None:
133
  texts = [""]
134
 
@@ -645,7 +646,7 @@ for epoch in range(start_epoch, start_epoch + num_epochs):
645
  if not fbp:
646
  if accelerator.sync_gradients:
647
  grad_val = accelerator.clip_grad_norm_(unet.parameters(), clip_grad_norm)
648
- grad = float(grad_val)
649
  optimizer.step()
650
  lr_scheduler.step()
651
  optimizer.zero_grad(set_to_none=True)
 
1
  #from comet_ml import Experiment
2
  import os
3
+ os.environ["NCCL_P2P_DISABLE"] = "1"
4
+ os.environ["NCCL_IB_DISABLE"] = "1"
5
+ os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
6
  import math
7
  import torch
8
  import numpy as np
 
12
  from torch.optim.lr_scheduler import LambdaLR
13
  from collections import defaultdict
14
  from diffusers import UNet2DConditionModel, AutoencoderKL,AutoencoderKLFlux2,AsymmetricAutoencoderKL,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
 
32
  # --------------------------- Параметры ---------------------------
33
  ds_path = "/workspace/sdxs-08b/datasets/d123_640_sd15"
34
  project = "unet"
35
+ ## total batch (split // num `GPU)
36
+ batch_size = 20
37
  base_learning_rate = 2e-6
38
+ min_learning_rate = 7e-7
39
+ num_epochs = 2
40
+ sample_interval_share = 10
41
  cfg_dropout = 0.10
42
  max_length = 248
43
  use_wandb = True
 
59
  torch.backends.cuda.enable_flash_sdp(True)
60
  torch.backends.cuda.enable_mem_efficient_sdp(True)
61
  torch.backends.cuda.enable_math_sdp(False) # Отключаем медленный вариант
 
62
  save_barrier = 1.05
63
  warmup_percent = 0.03
64
  percentile_clipping = 95
 
67
  clip_grad_norm = 1.0
68
  limit = 0
69
  checkpoints_folder = ""
 
70
  gradient_accumulation_steps = 1
71
+ dtype = torch.float32
72
+ mixed_precision = "no"
 
 
 
 
 
 
 
73
 
74
  # Параметры для диффузии
75
  n_diffusion_steps = 40
 
90
  if torch.cuda.is_available():
91
  torch.cuda.manual_seed_all(seed)
92
 
93
+ accelerator = Accelerator(
94
+ mixed_precision=mixed_precision,
95
+ gradient_accumulation_steps=gradient_accumulation_steps
96
+ )
97
+ device = accelerator.device
98
+
99
  print("init")
100
 
101
  # --------------------------- Инициализация WandB ---------------------------
 
129
  text_model = AutoModel.from_pretrained("text_encoder").to(device).eval()
130
  scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained("scheduler")
131
 
132
+ def encode_texts(texts, max_length=max_length):
133
  if texts is None:
134
  texts = [""]
135
 
 
646
  if not fbp:
647
  if accelerator.sync_gradients:
648
  grad_val = accelerator.clip_grad_norm_(unet.parameters(), clip_grad_norm)
649
+ grad = grad_val.float().item() if torch.is_tensor(grad_val) else float(grad_val)
650
  optimizer.step()
651
  lr_scheduler.step()
652
  optimizer.zero_grad(set_to_none=True)
unet/diffusion_pytorch_model.safetensors CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:9f32520c8e187f4a57efc07d926f8d2c0078a2e2e80adf5eaf811311aedbecd8
3
- size 1719263584
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:866fefbc206ff908dbc2fd5da4bcc21215fed0e3470c16b494eb4fb3d5810804
3
+ size 3438444088
unet1b.ipynb ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:37f536648a3fac58f603ac638ce41c79533c9e9f0da2fd10b6dc1bf8e35470e1
3
+ size 64797