@@ -73,5 +73,6 @@ def load_hydit(model_path, model_conf):
|
|||||||
load_device = load_device,
|
load_device = load_device,
|
||||||
offload_device = offload_device,
|
offload_device = offload_device,
|
||||||
current_device = "cpu",
|
current_device = "cpu",
|
||||||
|
size = 6 * (1024**3),
|
||||||
)
|
)
|
||||||
return model_patcher
|
return model_patcher
|
||||||
|
|||||||
@@ -97,6 +97,7 @@ class HYDiTTextEncode:
|
|||||||
|
|
||||||
def encode(self, text, text_t5, CLIP, T5):
|
def encode(self, text, text_t5, CLIP, T5):
|
||||||
# T5
|
# T5
|
||||||
|
T5.load_model()
|
||||||
t5_pre = T5.tokenizer(
|
t5_pre = T5.tokenizer(
|
||||||
text,
|
text,
|
||||||
max_length = T5.cond_stage_model.max_length,
|
max_length = T5.cond_stage_model.max_length,
|
||||||
@@ -117,6 +118,7 @@ class HYDiTTextEncode:
|
|||||||
t5_embs = t5_outs["hidden_states"][-1].float().cpu()
|
t5_embs = t5_outs["hidden_states"][-1].float().cpu()
|
||||||
|
|
||||||
# "clip"
|
# "clip"
|
||||||
|
CLIP.load_model()
|
||||||
clip_pre = CLIP.tokenizer(
|
clip_pre = CLIP.tokenizer(
|
||||||
text,
|
text,
|
||||||
max_length = CLIP.cond_stage_model.max_length,
|
max_length = CLIP.cond_stage_model.max_length,
|
||||||
|
|||||||
+14
-11
@@ -35,7 +35,7 @@ class mT5Model(torch.nn.Module):
|
|||||||
return self.transformer.load_state_dict(sd, strict=False)
|
return self.transformer.load_state_dict(sd, strict=False)
|
||||||
|
|
||||||
def to(self, *args, **kwargs):
|
def to(self, *args, **kwargs):
|
||||||
self.transformer.to(*args, **kwargs)
|
return self.transformer.to(*args, **kwargs)
|
||||||
|
|
||||||
class hyCLIPModel(torch.nn.Module):
|
class hyCLIPModel(torch.nn.Module):
|
||||||
def __init__(self, textmodel_json_config=None, device="cpu", max_length=77, freeze=True, dtype=None):
|
def __init__(self, textmodel_json_config=None, device="cpu", max_length=77, freeze=True, dtype=None):
|
||||||
@@ -64,47 +64,50 @@ class hyCLIPModel(torch.nn.Module):
|
|||||||
return self.transformer.load_state_dict(sd, strict=False)
|
return self.transformer.load_state_dict(sd, strict=False)
|
||||||
|
|
||||||
def to(self, *args, **kwargs):
|
def to(self, *args, **kwargs):
|
||||||
self.transformer.to(*args, **kwargs)
|
return self.transformer.to(*args, **kwargs)
|
||||||
|
|
||||||
class EXM_HyDiT_Tenc_Temp:
|
class EXM_HyDiT_Tenc_Temp:
|
||||||
def __init__(self, no_init=False, device="cpu", dtype=None, model_class="mT5", *kwargs):
|
def __init__(self, no_init=False, device="cpu", dtype=None, model_class="mT5", *kwargs):
|
||||||
if no_init:
|
if no_init:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
size = 8 if model_class == "mT5" else 2
|
||||||
|
if dtype == torch.float32:
|
||||||
|
size *= 2
|
||||||
|
size *= (1024**3)
|
||||||
|
|
||||||
if device == "auto":
|
if device == "auto":
|
||||||
size = 0
|
|
||||||
self.load_device = model_management.text_encoder_device()
|
self.load_device = model_management.text_encoder_device()
|
||||||
self.offload_device = model_management.text_encoder_offload_device()
|
self.offload_device = model_management.text_encoder_offload_device()
|
||||||
self.init_device = "cpu"
|
self.init_device = "cpu"
|
||||||
elif device == "cpu":
|
elif device == "cpu":
|
||||||
size = 0
|
size = 0 # doesn't matter
|
||||||
self.load_device = "cpu"
|
self.load_device = "cpu"
|
||||||
self.offload_device = "cpu"
|
self.offload_device = "cpu"
|
||||||
self.init_device="cpu"
|
self.init_device="cpu"
|
||||||
elif device.startswith("cuda"):
|
elif device.startswith("cuda"):
|
||||||
print("Direct CUDA device override!\nVRAM will not be freed by default.")
|
print("Direct CUDA device override!\nVRAM will not be freed by default.")
|
||||||
size = 0
|
size = 0 # not used
|
||||||
self.load_device = device
|
self.load_device = device
|
||||||
self.offload_device = device
|
self.offload_device = device
|
||||||
self.init_device = device
|
self.init_device = device
|
||||||
else:
|
else:
|
||||||
size = 0
|
|
||||||
self.load_device = model_management.get_torch_device()
|
self.load_device = model_management.get_torch_device()
|
||||||
self.offload_device = "cpu"
|
self.offload_device = "cpu"
|
||||||
self.init_device="cpu"
|
self.init_device="cpu"
|
||||||
|
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.device = device
|
self.device = self.load_device
|
||||||
if model_class == "mT5":
|
if model_class == "mT5":
|
||||||
self.cond_stage_model = mT5Model(
|
self.cond_stage_model = mT5Model(
|
||||||
device = device,
|
device = self.load_device,
|
||||||
dtype = dtype,
|
dtype = self.dtype,
|
||||||
)
|
)
|
||||||
tokenizer_args = {"subfolder": "t2i/mt5"}
|
tokenizer_args = {"subfolder": "t2i/mt5"}
|
||||||
else:
|
else:
|
||||||
self.cond_stage_model = hyCLIPModel(
|
self.cond_stage_model = hyCLIPModel(
|
||||||
device = device,
|
device = self.load_device,
|
||||||
dtype = dtype,
|
dtype = self.dtype,
|
||||||
)
|
)
|
||||||
tokenizer_args = {"subfolder": "t2i/tokenizer",}
|
tokenizer_args = {"subfolder": "t2i/tokenizer",}
|
||||||
self.tokenizer = AutoTokenizer.from_pretrained(
|
self.tokenizer = AutoTokenizer.from_pretrained(
|
||||||
|
|||||||
Reference in New Issue
Block a user