import os, base64, tempfile from typing import Dict, List, Any from io import BytesIO import torch from transformers import AutoProcessor from models.qwen3_vl_transformers import Qwen3VLForConditionalGeneration from models.pipeline import generate_image, DEFAULT_TIMESTEPS def add_special_tokens(tokenizer): tokenizer.boi_token = "<|boi_token|>" tokenizer.bor_token = "<|bor_token|>" tokenizer.eor_token = "<|eor_token|>" tokenizer.bot_token = "<|bot_token|>" tokenizer.tms_token = "<|tms_token|>" def get_tokenizer(processor): from transformers import PreTrainedTokenizerBase if isinstance(processor, PreTrainedTokenizerBase): return processor return processor.tokenizer class EndpointHandler: def __init__(self, path=""): self.processor = AutoProcessor.from_pretrained(path) self.model = Qwen3VLForConditionalGeneration.from_pretrained( path, torch_dtype=torch.float32, device_map="cuda" ).eval() self.tokenizer = get_tokenizer(self.processor) add_special_tokens(self.tokenizer) def __call__(self, data: Dict[str, Any]) -> List[Dict[str, Any]]: inputs = data.get("inputs") params = data.get("parameters", {}) if not inputs: return [{"error": "Missing \'inputs\' (text prompt)"}] model_type = params.get("model_type", "full") ref_images_b64 = params.get("ref_images", []) ref_image_paths = [] tmp_files = [] try: for i, b64_str in enumerate(ref_images_b64): img_bytes = base64.b64decode(b64_str) tmp = tempfile.NamedTemporaryFile(suffix=f"_{i}.png", delete=False) tmp.write(img_bytes) tmp.close() ref_image_paths.append(tmp.name) tmp_files.append(tmp.name) if model_type == "full": num_inference_steps = 36 guidance_scale = params.get("guidance_scale", 5.0) shift = 3.0 timesteps_list = None scheduler_name = "default" extra_kwargs = {} else: num_inference_steps = 28 guidance_scale = 0.0 shift = 1.0 timesteps_list = DEFAULT_TIMESTEPS scheduler_name = "flash" extra_kwargs = { "noise_scale_start": params.get("noise_scale_start", 7.5), "noise_scale_end": params.get("noise_scale_end", 7.5), "noise_clip_std": params.get("noise_clip_std", 2.5), } image = generate_image( model=self.model, processor=self.processor, prompt=inputs, ref_image_paths=ref_image_paths, height=params.get("height", 2048), width=params.get("width", 2048), num_inference_steps=num_inference_steps, guidance_scale=guidance_scale, shift=shift, timesteps_list=timesteps_list, scheduler_name=scheduler_name, seed=params.get("seed", 32), keep_original_aspect=params.get("keep_original_aspect", True), **extra_kwargs, ) except Exception as e: return [{"error": str(e)}] finally: for f in tmp_files: os.remove(f) return image