Merge pull request #3 from smthemex/Pr

init
This commit is contained in:
smthemex
2026-04-14 20:04:06 +08:00
committed by GitHub
4 changed files with 26 additions and 10 deletions
+9 -5
View File
@@ -6,7 +6,7 @@
Update
-----
* Need 64RAM+8VRAM
* Support gguf now, use less memory / 支持gguf,内存占用更少,模型在hg或者夸克云
1.Installation
@@ -29,14 +29,16 @@ pip install -r requirements.txt
3.checkpoints
----
* transformers/vae/clip [links](https://huggingface.co/jdopensource/JoyAI-Image-Edit)
* or [aliyun](https://pan.quark.cn/s/e20f511c921c)
* transformers/gguf/vae/clip [links](https://huggingface.co/jdopensource/JoyAI-Image-Edit)
* or [夸克云](https://pan.quark.cn/s/e20f511c921c)
* or [hg](https://huggingface.co/smthem/JoyAI-Image-Edit-merge-dit-gguf)
```
├── ComfyUI/models/
| ├── diffusion_models/
| ├──joy_image_transformer.safetensors
| ├──joy_image_transformer.safetensors # optional
| ├── gguf/
| ├──joy_image_transformer-Q8_0.gguf # optional
| ├── vae/
| ├──Wan2.1_VAE.pth
| ├── clips
@@ -49,12 +51,14 @@ pip install -r requirements.txt
![](https://github.com/smthemex/ComfyUI_JoyAI_Image/blob/main/example_workflows/example.png)
![](https://github.com/smthemex/ComfyUI_JoyAI_Image/blob/main/example_workflows/example2.png)
![](https://github.com/smthemex/ComfyUI_JoyAI_Image/blob/main/example_workflows/example3.png)
* GGUF
![](https://github.com/smthemex/ComfyUI_JoyAI_Image/blob/main/example_workflows/example_q.png)
5.Citation
----
```
@article{JoyAI2023,}
@jd-opensource
```
Binary file not shown.

After

Width:  |  Height:  |  Size: 570 KiB

+2 -2
View File
@@ -129,11 +129,11 @@ def load_dit(cfg, device: torch.device) -> torch.nn.Module:
# Ensure consistent dtype
param_dtypes = {param.dtype for param in model.parameters()}
if len(param_dtypes) > 1:
if len(param_dtypes) > 1 and not use_gguf:
logger.warning(
f"Model has mixed dtypes: {param_dtypes}. Converting to {dtype}")
model = model.to(dtype)
model.use_gguf=use_gguf
return model.eval()
def load_gguf_checkpoint(gguf_checkpoint_path):
+15 -3
View File
@@ -30,6 +30,7 @@ class BlockGPUManager:
self._original_block_ref = None
self._num_groups = 0 # 总批次数
self._group_loaded: list[bool] = []
self.use_gguf = False
def setup_for_inference(self, transformer_model):
@@ -41,6 +42,7 @@ class BlockGPUManager:
def _collect_managed_modules(self, transformer_model):
self.submodule = []
self._original_model_ref = transformer_model
self.use_gguf = getattr(transformer_model, "use_gguf", False)
self._original_block_ref = transformer_model.double_blocks
self._num_groups = (len(self._original_block_ref) + self.block_group_size - 1) // self.block_group_size
@@ -64,10 +66,14 @@ class BlockGPUManager:
group = nn.ModuleList()
for layer in self._original_block_ref[start_idx:end_idx]:
# 深拷贝当前层
cpu_layer = copy.deepcopy(layer)
if not self.use_gguf:
layer = copy.deepcopy(layer)
else:
# 记录原始设备以便后续恢复
layer._original_device = next(layer.parameters()).device
# 移动到目标设备
cpu_layer.to(self.device)
group.append(cpu_layer)
layer.to(self.device)
group.append(layer)
self.managed_modules[group_index] = group
self._group_loaded[group_index] = True
@@ -78,6 +84,12 @@ class BlockGPUManager:
return
group = self.managed_modules[group_index]
if self.use_gguf:
# 将层移回原始设备
for layer in group:
if hasattr(layer, '_original_device'):
layer.to(layer._original_device)
delattr(layer, '_original_device')
self.managed_modules[group_index] = None
self._group_loaded[group_index] = False