Quick FP8 fix.
This commit is contained in:
+11
-3
@@ -26,6 +26,14 @@ def load_dit(model_path, model_conf):
|
|||||||
state_dict = state_dict.get("model", state_dict)
|
state_dict = state_dict.get("model", state_dict)
|
||||||
parameters = comfy.utils.calculate_parameters(state_dict)
|
parameters = comfy.utils.calculate_parameters(state_dict)
|
||||||
unet_dtype = model_management.unet_dtype(model_params=parameters)
|
unet_dtype = model_management.unet_dtype(model_params=parameters)
|
||||||
|
load_device = comfy.model_management.get_torch_device()
|
||||||
|
offload_device = comfy.model_management.unet_offload_device()
|
||||||
|
|
||||||
|
# ignore fp8/etc and use directly for now
|
||||||
|
manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device)
|
||||||
|
if manual_cast_dtype:
|
||||||
|
print(f"DiT: falling back to {manual_cast_dtype}")
|
||||||
|
unet_dtype = manual_cast_dtype
|
||||||
|
|
||||||
model_conf["unet_config"]["num_classes"] = state_dict["y_embedder.embedding_table.weight"].shape[0] - 1 # adj. for empty
|
model_conf["unet_config"]["num_classes"] = state_dict["y_embedder.embedding_table.weight"].shape[0] - 1 # adj. for empty
|
||||||
|
|
||||||
@@ -46,8 +54,8 @@ def load_dit(model_path, model_conf):
|
|||||||
|
|
||||||
model_patcher = comfy.model_patcher.ModelPatcher(
|
model_patcher = comfy.model_patcher.ModelPatcher(
|
||||||
model,
|
model,
|
||||||
load_device=comfy.model_management.get_torch_device(),
|
load_device = load_device,
|
||||||
offload_device=comfy.model_management.unet_offload_device(),
|
offload_device = offload_device,
|
||||||
current_device="cpu"
|
current_device = "cpu",
|
||||||
)
|
)
|
||||||
return model_patcher
|
return model_patcher
|
||||||
|
|||||||
+10
-2
@@ -38,6 +38,14 @@ def load_pixart(model_path, model_conf):
|
|||||||
|
|
||||||
parameters = comfy.utils.calculate_parameters(state_dict)
|
parameters = comfy.utils.calculate_parameters(state_dict)
|
||||||
unet_dtype = model_management.unet_dtype(model_params=parameters)
|
unet_dtype = model_management.unet_dtype(model_params=parameters)
|
||||||
|
load_device = comfy.model_management.get_torch_device()
|
||||||
|
offload_device = comfy.model_management.unet_offload_device()
|
||||||
|
|
||||||
|
# ignore fp8/etc and use directly for now
|
||||||
|
manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device)
|
||||||
|
if manual_cast_dtype:
|
||||||
|
print(f"PixArt: falling back to {manual_cast_dtype}")
|
||||||
|
unet_dtype = manual_cast_dtype
|
||||||
|
|
||||||
model_conf = EXM_PixArt(model_conf) # convert to object
|
model_conf = EXM_PixArt(model_conf) # convert to object
|
||||||
model = comfy.model_base.BaseModel(
|
model = comfy.model_base.BaseModel(
|
||||||
@@ -64,8 +72,8 @@ def load_pixart(model_path, model_conf):
|
|||||||
|
|
||||||
model_patcher = comfy.model_patcher.ModelPatcher(
|
model_patcher = comfy.model_patcher.ModelPatcher(
|
||||||
model,
|
model,
|
||||||
load_device = comfy.model_management.get_torch_device(),
|
load_device = load_device,
|
||||||
offload_device = comfy.model_management.unet_offload_device(),
|
offload_device = offload_device,
|
||||||
current_device = "cpu",
|
current_device = "cpu",
|
||||||
)
|
)
|
||||||
return model_patcher
|
return model_patcher
|
||||||
|
|||||||
Reference in New Issue
Block a user