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