Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion modules/dataLoader/StableDiffusionFineTuneVaeDataLoader.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,7 +55,7 @@ def _setup_cache_device(
temp_device: torch.device,
config: TrainConfig,
):
model.to(self.temp_device)
model.release()

model.vae_to(train_device)

Expand Down
4 changes: 2 additions & 2 deletions modules/dataLoader/WuerstchenBaseDataLoader.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,7 @@ def _cache_modules(self, config: TrainConfig, model: WuerstchenModel, model_setu
]

def before_cache_image_fun():
model.to(self.temp_device)
model.release()
model.effnet_encoder_to(self.train_device)
model.eval()
torch_gc()
Expand Down Expand Up @@ -109,7 +109,7 @@ def _output_modules(self, config: TrainConfig, model: WuerstchenModel, model_set
output_names.append('pooled_text_encoder_output')

def before_cache_image_fun():
model.to(self.temp_device)
model.release()
model.effnet_encoder_to(self.train_device)
model.eval()
torch_gc()
Expand Down
4 changes: 2 additions & 2 deletions modules/dataLoader/mixin/DataLoaderText2ImageMixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -274,7 +274,7 @@ def _output_modules_from_out_names(
):
if before_cache_image_fun is None:
def prepare_vae():
model.to(self.temp_device)
model.release()
model.vae_to(self.train_device)
model.eval()
torch_gc()
Expand Down Expand Up @@ -340,7 +340,7 @@ def _cache_modules_from_names(

if before_cache_image_fun is None:
def prepare_vae():
model.to(self.temp_device)
model.release()
model.vae_to(self.train_device)
model.eval()
torch_gc()
Expand Down
5 changes: 4 additions & 1 deletion modules/model/BaseModel.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,8 +94,11 @@ def __init__(
self.autocast_context = nullcontext()
self.train_dtype = DataType.FLOAT_32

#park the whole model on the temp device to free VRAM. Models with on-demand components
#(which cannot be parked, only discarded and rebuilt) override this to free those components
#instead of moving them.
@abstractmethod
def to(self, device: torch.device):
def release(self):
pass

@abstractmethod
Expand Down
15 changes: 7 additions & 8 deletions modules/model/ChromaModel.py
Original file line number Diff line number Diff line change
Expand Up @@ -113,8 +113,7 @@ def vae_to(self, device: torch.device):

def text_encoder_to(self, device: torch.device):
if self.text_encoder is not None:
if self.text_encoder_offload_conductor is not None and \
self.text_encoder_offload_conductor.layer_offload_activated():
if self.text_encoder_offload_conductor is not None:
self.text_encoder_offload_conductor.to(device)
else:
self.text_encoder.to(device=device)
Expand All @@ -123,19 +122,19 @@ def text_encoder_to(self, device: torch.device):
self.text_encoder_lora.to(device)

def transformer_to(self, device: torch.device):
if self.transformer_offload_conductor is not None and \
self.transformer_offload_conductor.layer_offload_activated():
if self.transformer_offload_conductor is not None:
self.transformer_offload_conductor.to(device)
else:
self.transformer.to(device=device)

if self.transformer_lora is not None:
self.transformer_lora.to(device)

def to(self, device: torch.device):
self.vae_to(device)
self.text_encoder_to(device)
self.transformer_to(device)
def release(self):
temp_device = torch.device(self.train_config.temp_device)
self.vae_to(temp_device)
self.text_encoder_to(temp_device)
self.transformer_to(temp_device)

def eval(self):
self.vae.eval()
Expand Down
15 changes: 7 additions & 8 deletions modules/model/ErnieModel.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,26 +72,25 @@ def vae_to(self, device: torch.device):

def text_encoder_to(self, device: torch.device):
if self.text_encoder is not None:
if self.text_encoder_offload_conductor is not None and \
self.text_encoder_offload_conductor.layer_offload_activated():
if self.text_encoder_offload_conductor is not None:
self.text_encoder_offload_conductor.to(device)
else:
self.text_encoder.to(device=device)

def transformer_to(self, device: torch.device):
if self.transformer_offload_conductor is not None and \
self.transformer_offload_conductor.layer_offload_activated():
if self.transformer_offload_conductor is not None:
self.transformer_offload_conductor.to(device)
else:
self.transformer.to(device=device)

if self.transformer_lora is not None:
self.transformer_lora.to(device)

def to(self, device: torch.device):
self.vae_to(device)
self.text_encoder_to(device)
self.transformer_to(device)
def release(self):
temp_device = torch.device(self.train_config.temp_device)
self.vae_to(temp_device)
self.text_encoder_to(temp_device)
self.transformer_to(temp_device)

def eval(self):
self.vae.eval()
Expand Down
15 changes: 7 additions & 8 deletions modules/model/Flux2Model.py
Original file line number Diff line number Diff line change
Expand Up @@ -122,26 +122,25 @@ def vae_to(self, device: torch.device):

def text_encoder_to(self, device: torch.device):
if self.text_encoder is not None:
if self.text_encoder_offload_conductor is not None and \
self.text_encoder_offload_conductor.layer_offload_activated():
if self.text_encoder_offload_conductor is not None:
self.text_encoder_offload_conductor.to(device)
else:
self.text_encoder.to(device=device)

def transformer_to(self, device: torch.device):
if self.transformer_offload_conductor is not None and \
self.transformer_offload_conductor.layer_offload_activated():
if self.transformer_offload_conductor is not None:
self.transformer_offload_conductor.to(device)
else:
self.transformer.to(device=device)

if self.transformer_lora is not None:
self.transformer_lora.to(device)

def to(self, device: torch.device):
self.vae_to(device)
self.text_encoder_to(device)
self.transformer_to(device)
def release(self):
temp_device = torch.device(self.train_config.temp_device)
self.vae_to(temp_device)
self.text_encoder_to(temp_device)
self.transformer_to(temp_device)

def eval(self):
self.vae.eval()
Expand Down
15 changes: 7 additions & 8 deletions modules/model/FluxModel.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,8 +149,7 @@ def text_encoder_1_to(self, device: torch.device):

def text_encoder_2_to(self, device: torch.device):
if self.text_encoder_2 is not None:
if self.text_encoder_2_offload_conductor is not None and \
self.text_encoder_2_offload_conductor.layer_offload_activated():
if self.text_encoder_2_offload_conductor is not None:
self.text_encoder_2_offload_conductor.to(device)
else:
self.text_encoder_2.to(device=device)
Expand All @@ -159,19 +158,19 @@ def text_encoder_2_to(self, device: torch.device):
self.text_encoder_2_lora.to(device)

def transformer_to(self, device: torch.device):
if self.transformer_offload_conductor is not None and \
self.transformer_offload_conductor.layer_offload_activated():
if self.transformer_offload_conductor is not None:
self.transformer_offload_conductor.to(device)
else:
self.transformer.to(device=device)

if self.transformer_lora is not None:
self.transformer_lora.to(device)

def to(self, device: torch.device):
self.vae_to(device)
self.text_encoder_to(device)
self.transformer_to(device)
def release(self):
temp_device = torch.device(self.train_config.temp_device)
self.vae_to(temp_device)
self.text_encoder_to(temp_device)
self.transformer_to(temp_device)

def eval(self):
self.vae.eval()
Expand Down
18 changes: 8 additions & 10 deletions modules/model/HiDreamModel.py
Original file line number Diff line number Diff line change
Expand Up @@ -220,8 +220,7 @@ def text_encoder_2_to(self, device: torch.device):

def text_encoder_3_to(self, device: torch.device):
if self.text_encoder_3 is not None:
if self.text_encoder_3_offload_conductor is not None and \
self.text_encoder_3_offload_conductor.layer_offload_activated():
if self.text_encoder_3_offload_conductor is not None:
self.text_encoder_3_offload_conductor.to(device)
else:
self.text_encoder_3.to(device=device)
Expand All @@ -231,8 +230,7 @@ def text_encoder_3_to(self, device: torch.device):

def text_encoder_4_to(self, device: torch.device):
if self.text_encoder_4 is not None:
if self.text_encoder_4_offload_conductor is not None and \
self.text_encoder_4_offload_conductor.layer_offload_activated():
if self.text_encoder_4_offload_conductor is not None:
self.text_encoder_4_offload_conductor.to(device)
else:
self.text_encoder_4.to(device=device)
Expand All @@ -241,19 +239,19 @@ def text_encoder_4_to(self, device: torch.device):
self.text_encoder_4_lora.to(device)

def transformer_to(self, device: torch.device):
if self.transformer_offload_conductor is not None and \
self.transformer_offload_conductor.layer_offload_activated():
if self.transformer_offload_conductor is not None:
self.transformer_offload_conductor.to(device)
else:
self.transformer.to(device=device)

if self.transformer_lora is not None:
self.transformer_lora.to(device)

def to(self, device: torch.device):
self.vae_to(device)
self.text_encoder_to(device)
self.transformer_to(device)
def release(self):
temp_device = torch.device(self.train_config.temp_device)
self.vae_to(temp_device)
self.text_encoder_to(temp_device)
self.transformer_to(temp_device)

def eval(self):
self.vae.eval()
Expand Down
15 changes: 7 additions & 8 deletions modules/model/HunyuanVideoModel.py
Original file line number Diff line number Diff line change
Expand Up @@ -157,8 +157,7 @@ def text_encoder_to(self, device: torch.device):

def text_encoder_1_to(self, device: torch.device):
if self.text_encoder_1 is not None:
if self.text_encoder_1_offload_conductor is not None and \
self.text_encoder_1_offload_conductor.layer_offload_activated():
if self.text_encoder_1_offload_conductor is not None:
self.text_encoder_1_offload_conductor.to(device)
else:
self.text_encoder_1.to(device=device)
Expand All @@ -174,19 +173,19 @@ def text_encoder_2_to(self, device: torch.device):
self.text_encoder_2_lora.to(device)

def transformer_to(self, device: torch.device):
if self.transformer_offload_conductor is not None and \
self.transformer_offload_conductor.layer_offload_activated():
if self.transformer_offload_conductor is not None:
self.transformer_offload_conductor.to(device)
else:
self.transformer.to(device=device)

if self.transformer_lora is not None:
self.transformer_lora.to(device)

def to(self, device: torch.device):
self.vae_to(device)
self.text_encoder_to(device)
self.transformer_to(device)
def release(self):
temp_device = torch.device(self.train_config.temp_device)
self.vae_to(temp_device)
self.text_encoder_to(temp_device)
self.transformer_to(temp_device)

def eval(self):
self.vae.eval()
Expand Down
15 changes: 7 additions & 8 deletions modules/model/PixArtAlphaModel.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,8 +114,7 @@ def vae_to(self, device: torch.device):
self.vae.to(device=device)

def text_encoder_to(self, device: torch.device):
if self.text_encoder_offload_conductor is not None and \
self.text_encoder_offload_conductor.layer_offload_activated():
if self.text_encoder_offload_conductor is not None:
self.text_encoder_offload_conductor.to(device)
else:
self.text_encoder.to(device=device)
Expand All @@ -124,19 +123,19 @@ def text_encoder_to(self, device: torch.device):
self.text_encoder_lora.to(device)

def transformer_to(self, device: torch.device):
if self.transformer_offload_conductor is not None and \
self.transformer_offload_conductor.layer_offload_activated():
if self.transformer_offload_conductor is not None:
self.transformer_offload_conductor.to(device)
else:
self.transformer.to(device=device)

if self.transformer_lora is not None:
self.transformer_lora.to(device)

def to(self, device: torch.device):
self.vae_to(device)
self.text_encoder_to(device)
self.transformer_to(device)
def release(self):
temp_device = torch.device(self.train_config.temp_device)
self.vae_to(temp_device)
self.text_encoder_to(temp_device)
self.transformer_to(temp_device)

def eval(self):
self.vae.eval()
Expand Down
15 changes: 7 additions & 8 deletions modules/model/QwenModel.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,8 +81,7 @@ def vae_to(self, device: torch.device):

def text_encoder_to(self, device: torch.device): #TODO share more code between models
if self.text_encoder is not None:
if self.text_encoder_offload_conductor is not None and \
self.text_encoder_offload_conductor.layer_offload_activated():
if self.text_encoder_offload_conductor is not None:
self.text_encoder_offload_conductor.to(device)
else:
self.text_encoder.to(device=device)
Expand All @@ -91,19 +90,19 @@ def text_encoder_to(self, device: torch.device): #TODO share more code between m
self.text_encoder_lora.to(device)

def transformer_to(self, device: torch.device):
if self.transformer_offload_conductor is not None and \
self.transformer_offload_conductor.layer_offload_activated():
if self.transformer_offload_conductor is not None:
self.transformer_offload_conductor.to(device)
else:
self.transformer.to(device=device)

if self.transformer_lora is not None:
self.transformer_lora.to(device)

def to(self, device: torch.device):
self.vae_to(device)
self.text_encoder_to(device)
self.transformer_to(device)
def release(self):
temp_device = torch.device(self.train_config.temp_device)
self.vae_to(temp_device)
self.text_encoder_to(temp_device)
self.transformer_to(temp_device)

def eval(self):
self.vae.eval()
Expand Down
15 changes: 7 additions & 8 deletions modules/model/SanaModel.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,8 +116,7 @@ def vae_to(self, device: torch.device):
self.vae.to(device=device)

def text_encoder_to(self, device: torch.device):
if self.text_encoder_offload_conductor is not None and \
self.text_encoder_offload_conductor.layer_offload_activated():
if self.text_encoder_offload_conductor is not None:
self.text_encoder_offload_conductor.to(device)
else:
self.text_encoder.to(device=device)
Expand All @@ -126,19 +125,19 @@ def text_encoder_to(self, device: torch.device):
self.text_encoder_lora.to(device)

def transformer_to(self, device: torch.device):
if self.transformer_offload_conductor is not None and \
self.transformer_offload_conductor.layer_offload_activated():
if self.transformer_offload_conductor is not None:
self.transformer_offload_conductor.to(device)
else:
self.transformer.to(device=device)

if self.transformer_lora is not None:
self.transformer_lora.to(device)

def to(self, device: torch.device):
self.vae_to(device)
self.text_encoder_to(device)
self.transformer_to(device)
def release(self):
temp_device = torch.device(self.train_config.temp_device)
self.vae_to(temp_device)
self.text_encoder_to(temp_device)
self.transformer_to(temp_device)

def eval(self):
self.vae.eval()
Expand Down
Loading