geonhow/abstract-sd
0
1import os2import torch3import base644from diffusers import StableDiffusionXLPipeline, DiffusionPipeline5from PIL import Image6from flask import Flask, request, jsonify7from PIL import Image8from io import BytesIO9 10# Flask 앱 생성11app = Flask(__name__)12 13# 캐시 디렉토리 설정14os.environ["TRANSFORMERS_CACHE"] = "/app/.cache/huggingface"15os.makedirs(os.environ["TRANSFORMERS_CACHE"], exist_ok=True)16 17# CUDA 또는 CPU 설정18device = "cuda" if torch.cuda.is_available() else "cpu"19print(f"Using device: {device}")20 21# ✅ 모델 로드22model_id = "geonhow/abstract_finetuned_model_initial" # Fine-Tuned SDXL 모델23 24dtype = torch.float16 if device == "cuda" else torch.float3225pipeline = StableDiffusionXLPipeline.from_pretrained(26 model_id,27 torch_dtype=dtype28).to(device)29from diffusers import DPMSolverMultistepScheduler30pipeline.scheduler = DPMSolverMultistepScheduler.from_config(pipeline.scheduler.config)31# 2. Refiner 로드 (text_encoder_2, vae 공유)32refiner = DiffusionPipeline.from_pretrained(33 "stabilityai/stable-diffusion-xl-refiner-1.0",34 text_encoder_2=pipeline.text_encoder_2,35 vae=pipeline.vae,36 torch_dtype=dtype,37 use_safetensors=True,38).to(device)39 40# 6. 생성 파라미터41num_inference_steps = 2542high_noise_frac = 0.98 # base가 80%, refiner가 20%43 44 45# ✅ 이미지 Base64 변환 함수46def encode_image(image):47 buffered = BytesIO()48 image.save(buffered, format="PNG")49 return base64.b64encode(buffered.getvalue()).decode()50 51# ✅ API 엔드포인트: 감정 기반 이미지 생성52@app.route("/generate", methods=["POST"])53def generate_image():54 data = request.get_json() or {}55 prompt = data.get("prompt", "")56 prompt_2 = data.get("prompt_2","")57 58 if not prompt:59 return jsonify({"error": "Prompt is required!"}), 40060 61 # 1단계: base 모델로 latent 생성62 with torch.no_grad():63 image = pipeline(64 prompt=prompt,65 prompt_2=prompt_2,66 num_inference_steps=num_inference_steps,67 guidance_scale=15,68 ).images[0]69 70 71 72 # 중앙 크롭 (1000x1000)73 crop_size = 100074 width, height = image.size75 left = max((width - crop_size) // 2, 0)76 top = max((height - crop_size) // 2, 0)77 right = left + crop_size78 bottom = top + crop_size79 cropped_image = image.crop((left, top, right, bottom))80 81 return jsonify({"image": encode_image(cropped_image)})82 83# 기본 실행84if __name__ == "__main__":85 app.run(host="0.0.0.0", port=7860, debug=True)