Singularity666/Magix
1
1import gradio as gr2import os3import shutil4from main import fine_tune_model5from diffusers import StableDiffusionPipeline, DDIMScheduler6import torch7 8MODEL_NAME = "runwayml/stable-diffusion-v1-5"9OUTPUT_DIR = "/home/user/app/stable_diffusion_weights/custom_model"10 11def fine_tune(instance_prompt, image1, image2=None):12 instance_data_dir = "/home/user/app/instance_images"13 14 try:15 if os.path.exists(instance_data_dir):16 shutil.rmtree(instance_data_dir)17 os.makedirs(instance_data_dir, exist_ok=True)18 19 image1.save(os.path.join(instance_data_dir, "instance_0.png"))20 if image2 is not None:21 image2.save(os.path.join(instance_data_dir, "instance_1.png"))22 23 fine_tune_model(instance_data_dir, instance_prompt, MODEL_NAME, OUTPUT_DIR)24 return "Model fine-tuning complete."25 except Exception as e:26 return str(e)27 28def generate_images(prompt, num_samples, height, width, num_inference_steps, guidance_scale):29 try:30 if not os.path.exists(OUTPUT_DIR):31 return "The model path does not exist."32 33 pipe = StableDiffusionPipeline.from_pretrained(OUTPUT_DIR, safety_checker=None, torch_dtype=torch.float16).to("cuda")34 pipe.scheduler = DDIMScheduler.from_config(pipe.scheduler.config)35 g_cuda = torch.Generator(device='cuda').manual_seed(1337)36 37 with torch.autocast("cuda"), torch.inference_mode():38 images = pipe(39 prompt, height=height, width=width, num_images_per_prompt=num_samples,40 num_inference_steps=num_inference_steps, guidance_scale=guidance_scale, generator=g_cuda41 ).images42 43 return images44 except Exception as e:45 return str(e)46 47def gradio_app():48 with gr.Blocks() as demo:49 with gr.Tab("Fine-Tune Model"):50 with gr.Row():51 with gr.Column():52 instance_prompt = gr.Textbox(label="Instance Prompt")53 image1 = gr.Image(label="Upload Image 1", type="pil")54 image2 = gr.Image(label="Upload Image 2 (Optional)", type="pil")55 fine_tune_button = gr.Button("Fine-Tune Model")56 output_text = gr.Textbox(label="Output")57 fine_tune_button.click(fine_tune, inputs=[instance_prompt, image1, image2], outputs=output_text)58 59 with gr.Tab("Generate Images"):60 with gr.Row():61 with gr.Column():62 prompt = gr.Textbox(label="Prompt")63 num_samples = gr.Number(label="Number of Samples", value=1)64 guidance_scale = gr.Number(label="Guidance Scale", value=7.5)65 height = gr.Number(label="Height", value=512)66 width = gr.Number(label="Width", value=512)67 num_inference_steps = gr.Slider(label="Steps", value=50, minimum=1, maximum=100)68 generate_button = gr.Button("Generate Images")69 with gr.Column():70 gallery = gr.Gallery(label="Generated Images")71 generate_button.click(generate_images, inputs=[prompt, num_samples, height, width, num_inference_steps, guidance_scale], outputs=gallery)72 73 demo.launch()74 75if __name__ == "__main__":76 gradio_app()