Team Ai
Apppublic

bala1802/StableDiffusionModel

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
app.py41 linesDownload Raw Back to root
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)