File size: 5,425 Bytes
339703a
73d18ae
339703a
73d18ae
339703a
 
73d18ae
 
 
339703a
73d18ae
 
c2860de
 
73d18ae
 
 
 
 
 
 
 
 
 
 
c2860de
 
73d18ae
 
 
 
 
 
 
 
c2860de
73d18ae
c2860de
73d18ae
 
c2860de
73d18ae
 
 
 
 
c2860de
73d18ae
 
 
 
 
 
ed6cd4c
73d18ae
 
 
 
c2860de
73d18ae
 
 
 
 
 
 
 
6798543
73d18ae
 
6798543
73d18ae
 
 
 
 
 
 
 
6798543
73d18ae
 
 
 
 
 
 
 
 
e93e938
73d18ae
 
 
 
 
 
 
 
 
 
 
c2860de
73d18ae
 
6798543
73d18ae
 
6798543
73d18ae
 
 
 
 
6798543
73d18ae
6798543
73d18ae
 
 
 
 
c2860de
73d18ae
 
 
 
c2860de
 
 
 
73d18ae
 
 
 
339703a
14c15ad
73d18ae
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
339703a
 
 
73d18ae
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
import os
import spaces
import gc
import shutil
import gradio as gr
import torch
import safetensors.torch
from huggingface_hub import hf_hub_download, HfApi, login
from accelerate import init_empty_weights

# --- Imports from your local modules ---
# Ensure the folder 'qwenimage' is present in the root directory
from qwenimage.transformer_qwenimage import QwenImageTransformer2DModel

# Configuration for the specific Qwen Transformer
TRANSFORMER_CONFIG = {
    "attention_head_dim": 128,
    "axes_dims_rope": [16, 56, 56],
    "guidance_embeds": False,
    "in_channels": 64,
    "joint_attention_dim": 3584,
    "num_attention_heads": 24,
    "num_layers": 60,
    "out_channels": 16,
    "patch_size": 2
}

@spaces.GPU(duration=300)
def convert_and_upload(hf_token, target_repo_id, private_repo):
    """
    Downloads raw weights, converts keys, saves locally, and uploads to HF.
    """
    local_dir = "converted_qwen_transformer"
    source_repo = "Phr00t/Qwen-Image-Edit-Rapid-AIO"
    source_filename = "v19/Qwen-Rapid-AIO-NSFW-v19.safetensors"
    
    yield f"๐Ÿš€ Starting process...\nAuthenticating with Hugging Face..."
    
    if not hf_token:
        raise gr.Error("Please provide a Write-enabled Hugging Face Token.")
    
    try:
        login(token=hf_token)
        api = HfApi(token=hf_token)
    except Exception as e:
        raise gr.Error(f"Authentication failed: {e}")

    # 1. Download
    yield f"๐Ÿ“ฅ Downloading {source_filename} from {source_repo}..."
    try:
        checkpoint_path = hf_hub_download(repo_id=source_repo, filename=source_filename)
    except Exception as e:
        raise gr.Error(f"Download failed: {e}")

    # 2. Initialize Empty Model
    yield "๐Ÿ—๏ธ Initializing empty model architecture..."
    with init_empty_weights():
        model = QwenImageTransformer2DModel(**TRANSFORMER_CONFIG)

    # 3. Load and Filter Keys
    yield "๐Ÿ”‘ Loading state dict and filtering keys (removing 'model.diffusion_model.')..."
    try:
        state_dict = safetensors.torch.load_file(checkpoint_path, device="cpu")
        
        new_state_dict = {}
        prefix = "model.diffusion_model."
        ignored_keys = ["__index_timestep_zero__", "iteration", "global_step"]

        for key, value in state_dict.items():
            if key in ignored_keys:
                continue
            if key.startswith(prefix):
                new_key = key[len(prefix):]
                new_state_dict[new_key] = value
        
        del state_dict
        gc.collect()
    except Exception as e:
        raise gr.Error(f"Error processing keys: {e}")

    # 4. Load Weights into Model
    yield "โš–๏ธ Loading weights into the model object..."
    try:
        # assign=True is needed for accelerate's init_empty_weights
        model.load_state_dict(new_state_dict, assign=True, strict=False)
        del new_state_dict
        gc.collect()
    except Exception as e:
        raise gr.Error(f"Error loading weights into model: {e}")

    # 5. Save Locally
    if os.path.exists(local_dir):
        shutil.rmtree(local_dir)
    os.makedirs(local_dir, exist_ok=True)
    
    yield f"๐Ÿ’พ Saving converted model to local directory: {local_dir}..."
    try:
        # This saves both config.json and diffusion_pytorch_model.safetensors
        model.save_pretrained(local_dir, safe_serialization=True)
    except Exception as e:
        raise gr.Error(f"Error saving local model: {e}")

    # 6. Upload to Hugging Face
    yield f"โ˜๏ธ Uploading to Hugging Face Repo: {target_repo_id}..."
    try:
        # Create repo if it doesn't exist
        api.create_repo(repo_id=target_repo_id, private=private_repo, exist_ok=True)
        
        api.upload_folder(
            folder_path=local_dir,
            repo_id=target_repo_id,
            commit_message="Upload converted Qwen-Image-Edit Transformer"
        )
    except Exception as e:
        raise gr.Error(f"Upload failed: {e}")
        
    # Cleanup
    shutil.rmtree(local_dir)
    gc.collect()
    
    yield f"โœ… Success! Model uploaded to https://huggingface.co/{target_repo_id}"

# --- Gradio UI ---

css = """
#col-container { max_width: 700px; margin: 0 auto; }
"""

with gr.Blocks() as demo:
    with gr.Column(elem_id="col-container"):
        gr.Markdown("# ๐Ÿ”„ Qwen Transformer Converter & Uploader")
        gr.Markdown(
            "This tool downloads the raw checkpoints for `Qwen-Image-Edit`, extracts the transformer, "
            "fixes the key names, and uploads the clean `diffusers`-ready model to your Hugging Face account."
        )
        
        with gr.Group():
            hf_token = gr.Textbox(
                label="Hugging Face Token (Write Access)", 
                placeholder="hf_...", 
                type="password"
            )
            target_repo = gr.Textbox(
                label="Target Repository ID", 
                placeholder="username/my-converted-qwen-transformer"
            )
            is_private = gr.Checkbox(label="Make Repo Private", value=True)
            
        convert_btn = gr.Button("Convert & Upload", variant="primary")
        status_output = gr.Textbox(label="Status Log", interactive=False, lines=6)

    convert_btn.click(
        fn=convert_and_upload,
        inputs=[hf_token, target_repo, is_private],
        outputs=[status_output]
    )

if __name__ == "__main__":
    demo.queue().launch(theme=gr.themes.Soft(), css=css)