import gradio as gr import tempfile import os from abc import ABC, abstractmethod from typing import Dict import torch from safetensors.torch import load_file, save_file # ========================================== # 1. ИНТЕРФЕЙСЫ (ISP & DIP) # ========================================== class IKeyMapper(ABC): """ Что: Абстракция для правил переименования ключей. Почему: Позволяет добавлять новые форматы конвертации (OCP), не меняя ядро программы. """ @abstractmethod def map_key(self, key: str) -> str: pass class IModelIO(ABC): """ Что: Абстракция для чтения и записи весов. Почему: Отвязывает логику конвертации от конкретного формата (safetensors, pt) и файловой системы (DIP). """ @abstractmethod def load(self, path: str) -> Dict[str, torch.Tensor]: pass @abstractmethod def save(self, state_dict: Dict[str, torch.Tensor], path: str) -> None: pass # ========================================== # 2. БИЗНЕС-ЛОГИКА (SRP & OCP) # ========================================== class DiffusersToKohyaMapper(IKeyMapper): """ Что: Реализация маппера для перевода Diffusers -> Kohya (ComfyUI/A1111). Почему: Diffusers использует иерархию через точку (unet.up_blocks...), а ComfyUI ожидает плоскую структуру с нижними подчеркиваниями (lora_unet_up_blocks...). """ def __init__(self): # YAGNI: Указываем только те префиксы, которые реально нужны для SD/SDXL. self.prefix_map = { "unet": "lora_unet", "text_encoder": "lora_te1", "text_encoder_2": "lora_te2", } def map_key(self, key: str) -> str: parts = key.split('.') # Если префикс неизвестен, возвращаем как есть (защита от поломки) if not parts or parts[0] not in self.prefix_map: return key base = self.prefix_map[parts[0]] # Отделяем суффикс типа тензора (weight/alpha), чтобы не сломать его при замене точек suffix = "" if key.endswith(".weight"): suffix = ".weight" key = key[:-7] elif key.endswith(".alpha"): suffix = ".alpha" key = key[:-6] # Заменяем точки на подчеркивания в теле ключа # Пример: up_blocks.0.attentions.1.lora.up -> up_blocks_0_attentions_1_lora_up body = key[len(parts[0])+1:] body = body.replace(".lora.up", "_lora_up").replace(".lora.down", "_lora_down") body = body.replace(".", "_") return f"{base}_{body}{suffix}" class LoraConverter: """ Что: Оркестратор процесса конвертации. Почему: Делегирует чтение/запись и маппинг внедренным зависимостям (Dependency Injection). """ def __init__(self, mapper: IKeyMapper, io_handler: IModelIO): self._mapper = mapper self._io_handler = io_handler def convert_file(self, input_path: str, output_path: str) -> None: """ Что: Основной цикл конвертации словаря весов. Почему: Вынесено в отдельный метод для сохранения DRY, если понадобится вызывать это из API, а не только из UI. """ state_dict = self._io_handler.load(input_path) converted_dict = {} for key, tensor in state_dict.items(): new_key = self._mapper.map_key(key) converted_dict[new_key] = tensor self._io_handler.save(converted_dict, output_path) # ========================================== # 3. ИНФРАСТРУКТУРА # ========================================== class SafetensorsHandler(IModelIO): """ Что: Конкретная реализация работы с Safetensors. Почему: Инкапсулирует библиотеку safetensors. Если завтра потребуется поддержка .pt, мы создадим PtHandler, не трогая LoraConverter (LSP). """ def load(self, path: str) -> Dict[str, torch.Tensor]: return load_file(path) def save(self, state_dict: Dict[str, torch.Tensor], path: str) -> None: save_file(state_dict, path) # ========================================== # 4. СЛОЙ ПРЕДСТАВЛЕНИЯ (UI) # ========================================== def process_upload(file_obj) -> str: """ Что: Точка входа для Gradio. Composition Root. Почему: Здесь собираются воедино все слои приложения (KISS). """ if file_obj is None: return None # Сборка зависимостей (Dependency Injection) mapper = DiffusersToKohyaMapper() io_handler = SafetensorsHandler() converter = LoraConverter(mapper=mapper, io_handler=io_handler) # Создание временного файла для безопасного вывода output_fd, output_path = tempfile.mkstemp(suffix="_comfyui.safetensors") os.close(output_fd) # Запуск логики converter.convert_file(file_obj.name, output_path) return output_path # Настройка интерфейса Gradio with gr.Blocks(title="LoRA Diffusers to ComfyUI Converter") as app: gr.Markdown("## Конвертер архитектуры LoRA (Diffusers -> ComfyUI)") gr.Markdown("Загрузите LoRA в формате Diffusers (обычно выдает ошибку `lora key not loaded`), и получите файл, готовый к работе в ComfyUI.") with gr.Row(): file_input = gr.File(label="Загрузить LoRA (.safetensors)", file_types=[".safetensors"]) file_output = gr.File(label="Скачать конвертированную LoRA", interactive=False) convert_btn = gr.Button("Конвертировать", variant="primary") convert_btn.click(fn=process_upload, inputs=file_input, outputs=file_output) if __name__ == "__main__": app.launch()