From 5fcb69fee673289f843ef033657e33602c4ace52 Mon Sep 17 00:00:00 2001 From: taabata <57796911+taabata@users.noreply.github.com> Date: Sat, 28 Dec 2024 04:02:11 +0300 Subject: [PATCH] Add files via upload --- SANA/app.py | 153 +++++++++++++++++++++++++++------------------------- 1 file changed, 81 insertions(+), 72 deletions(-) diff --git a/SANA/app.py b/SANA/app.py index 66080ea..6000d6a 100644 --- a/SANA/app.py +++ b/SANA/app.py @@ -1,33 +1,19 @@ - - from flask import Flask, request from PIL import Image -import subprocess, json, os +import json, os import torch from pipes.sana_img2img import SanaPipelineImg2Img from pipes.sana_pag import SanaPAGPipeline from pipes.sana_pag_img2img import SanaPipelineImg2ImgPAG import numpy as np -import argparse from signal import SIGKILL -parser = argparse.ArgumentParser() - -parser.add_argument("--prompt","-p") -parser.add_argument("--steps","-s") -parser.add_argument("--cfg","-c") - -args = parser.parse_args() - -steps = args.steps -cfg = args.cfg if args.cfg else 5.0 -prompt = args.prompt params = { - "steps":steps, - "cfg":cfg, - "prompt":prompt + "steps":12, + "cfg":5.0, + "prompt":"" } @@ -50,70 +36,93 @@ def endApp(): @app.route("/getSharedData",methods=["POST","GET"]) def getSharedData(): - global flag, output_image + global flag, output_image, embeds return {"flag":flag,"image":output_image,"embeds":embeds} @app.route("/encode",methods=["POST","GET"]) def encode(): - global params - print("encoding......") - subprocess.Popen(["python3","getembeds.py","--prompt",request.json["prompt"],"--model",request.json["model"]]) + global params, embeds, flag + try: + model_name = request.json["model"] + pipe = SanaPipelineImg2Img.from_pretrained( + model_name, + torch_dtype=torch.bfloat16, + transformer=None, + vae=None + + ) + pipe.to('cpu') + prompt = request.json["prompt"] + negative_prompt = request.json["negative_prompt"] + embeds = pipe( + prompt=prompt, + negative_prompt=negative_prompt, + height=512, + width=512, + guidance_scale=5.0, + num_inference_steps=2, + textonly=True + ) + data = {} + for i in embeds: + embs = json.dumps(embeds[i].type(torch.float32).numpy().astype(np.float32).tolist()) + data[i] = embs + embeds = data + flag = False + except: + flag = False return {} -@app.route("/getEmbeds",methods=["POST","GET"]) -def getEmbeds(): - global embeds,flag - embeds = request.json - flag = False - return {} @app.route("/diffuse",methods=["POST","GET"]) def diffuse(): global embeds,params, output_image, flag - print("diffusing......") - device = request.json["device"] - if request.json["img2img"] == "enable": - pipe = SanaPipelineImg2ImgPAG.from_pretrained( - request.json["model"], - torch_dtype = torch.float16 if device=="cuda" else torch.float32, - text_encoder=None - ) - if device=="cuda": - pipe.enable_model_cpu_offload() - image = pipe( - prompt_embeds = torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_embeds"]))).half().to('cuda'), - prompt_attention_mask= torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_attention_mask"]))).half().to('cuda'), - negative_prompt_embeds=torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_embeds"]))).half().to('cuda'), - negative_prompt_attention_mask = torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_attention_mask"]))).half().to('cuda'), - height=int(request.json["height"]), - width=int(request.json["width"]), - guidance_scale=float(request.json["cfg"]), - pag_scale=float(request.json["pag_scale"]), - num_inference_steps=int(request.json["steps"]), - image=Image.fromarray(np.array(json.loads(request.json["image"]),dtype="uint8")), - strength=float(request.json["strength"]) - )[0] - else: - pipe = SanaPAGPipeline.from_pretrained( - request.json["model"], - torch_dtype = torch.float16 if device=="cuda" else torch.float32, - text_encoder=None - ) - if device=="cuda": - pipe.enable_model_cpu_offload() - image = pipe( - prompt_embeds = torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_embeds"]))).half().to('cuda'), - prompt_attention_mask= torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_attention_mask"]))).half().to('cuda'), - negative_prompt_embeds=torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_embeds"]))).half().to('cuda'), - negative_prompt_attention_mask = torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_attention_mask"]))).half().to('cuda'), - height=int(request.json["height"]), - width=int(request.json["width"]), - guidance_scale=float(request.json["cfg"]), - pag_scale=float(request.json["pag_scale"]), - num_inference_steps=int(request.json["steps"]), - )[0] - output_image = json.dumps(np.array(image[0]).tolist()) - flag = False + try: + device = request.json["device"] + if request.json["img2img"] == "enable": + pipe = SanaPipelineImg2ImgPAG.from_pretrained( + request.json["model"], + torch_dtype = torch.float16 if device=="cuda" else torch.float32, + text_encoder=None + ) + if device=="cuda": + pipe.enable_model_cpu_offload() + image = pipe( + prompt_embeds = torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_embeds"]))).half().to('cuda') if device=="cuda" else torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_embeds"]))), + prompt_attention_mask= torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_attention_mask"]))).half().to('cuda') if device=="cuda" else torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_attention_mask"]))), + negative_prompt_embeds=torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_embeds"]))).half().to('cuda') if device=="cuda" else torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_embeds"]))), + negative_prompt_attention_mask = torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_attention_mask"]))).half().to('cuda') if device=="cuda" else torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_attention_mask"]))), + height=int(request.json["height"]), + width=int(request.json["width"]), + guidance_scale=float(request.json["cfg"]), + pag_scale=float(request.json["pag_scale"]), + num_inference_steps=int(request.json["steps"]), + image=Image.fromarray(np.array(json.loads(request.json["image"]),dtype="uint8")), + strength=float(request.json["strength"]) + )[0] + else: + pipe = SanaPAGPipeline.from_pretrained( + request.json["model"], + torch_dtype = torch.float16 if device=="cuda" else torch.float32, + text_encoder=None + ) + if device=="cuda": + pipe.enable_model_cpu_offload() + image = pipe( + prompt_embeds = torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_embeds"]))).half().to('cuda') if device=="cuda" else torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_embeds"]))), + prompt_attention_mask= torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_attention_mask"]))).half().to('cuda') if device=="cuda" else torch.Tensor(np.array(json.loads(request.json["embeds"]["prompt_attention_mask"]))), + negative_prompt_embeds=torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_embeds"]))).half().to('cuda') if device=="cuda" else torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_embeds"]))), + negative_prompt_attention_mask = torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_attention_mask"]))).half().to('cuda') if device=="cuda" else torch.Tensor(np.array(json.loads(request.json["embeds"]["negative_prompt_attention_mask"]))), + height=int(request.json["height"]), + width=int(request.json["width"]), + guidance_scale=float(request.json["cfg"]), + pag_scale=float(request.json["pag_scale"]), + num_inference_steps=int(request.json["steps"]), + )[0] + output_image = json.dumps(np.array(image[0]).tolist()) + flag = False + except: + flag = False return{}