bala1802/StableDiffusionModel
0
1import torch2import gradio as gr3 4import prediction5import model6import diffusion_loss7 8device = 'cuda' if torch.cuda.is_available() else 'cpu'9 10pipe = model.initialize_diffusion_model()11 12def generate(prompt, loss_function=None):13 return prediction.predict(prompt=prompt, pipe=pipe, loss_function=loss_function)14 15def process_input(prompt, loss_function, button):16 if button:17 if loss_function is None or loss_function == "No Loss":18 return generate(prompt, loss_function=None)19 elif loss_function == "Blue Channel":20 return generate(prompt, loss_function=diffusion_loss.blue_channel)21 elif loss_function == "Saturation":22 return generate(prompt, loss_function=diffusion_loss.saturation)23 elif loss_function == "Elastic Deformation":24 return generate(prompt, loss_function=diffusion_loss.elastic_transform)25 else:26 return generate(prompt, loss_function=None)27 else:28 return None29 30iface = gr.Interface(31 fn=process_input,32 inputs=[33 gr.Textbox("prompt", label="Enter Prompt"),34 gr.Dropdown(["No Loss", "Blue Channel", "Saturation", 'Elastic Deformation'], label='Choose Augmentation'),35 gr.Button("Loss Function")],36 37 outputs = gr.Image(type="pil")38)39 40if __name__ == "__main__":41 iface.launch(show_api=False, share=True)