realistic-vision-v6-axera-hw / sample_inference.py
gregm123456's picture
Upload folder using huggingface_hub
758e1bd verified
Raw
History Blame Contribute Delete
3.9 kB
import os
import time
import argparse
import numpy as np
import torch
import axengine
from PIL import Image
from transformers import CLIPTokenizer
def get_args():
parser = argparse.ArgumentParser(description="Axera Realistic Vision V6.0 B1 Inference")
parser.add_argument("--prompt", type=str, default="A serene portrait of an elderly man with silver hair, warm smile, greenhouse background, highly detailed", help="Text prompt")
parser.add_argument("--output", type=str, default="output.png", help="Output image path")
return parser.parse_args()
def get_alphas_cumprod():
betas = torch.linspace(0.00085 ** 0.5, 0.012 ** 0.5, 1000, dtype=torch.float32) ** 2
alphas = 1.0 - betas
alphas_cumprod = torch.cumprod(alphas, dim=0).detach().numpy()
final_alphas_cumprod = alphas_cumprod[0]
return alphas_cumprod, final_alphas_cumprod
def main():
args = get_args()
# Paths
base_dir = os.path.dirname(__file__)
tokenizer_dir = os.path.join(base_dir, "tokenizer")
text_encoder_model = os.path.join(base_dir, "sd15_text_encoder_sim.axmodel")
unet_model = os.path.join(base_dir, "unet.axmodel")
vae_decoder_model = os.path.join(base_dir, "vae_decoder.axmodel")
time_input_path = os.path.join(base_dir, "time_input_txt2img.npy")
print(f"Loading models...")
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_dir)
text_encoder = axengine.InferenceSession(text_encoder_model)
unet_session = axengine.InferenceSession(unet_model)
vae_decoder = axengine.InferenceSession(vae_decoder_model)
time_embeddings = np.load(time_input_path)
alphas_cumprod, final_alphas_cumprod = get_alphas_cumprod()
timesteps = np.array([999, 759, 499, 259]).astype(np.int64)
# 1. Text Encoding
print(f"Encoding prompt: {args.prompt}")
text_inputs = tokenizer(args.prompt, padding="max_length", max_length=77, truncation=True, return_tensors="pt")
prompt_embeds = text_encoder.run(None, {"input_ids": text_inputs.input_ids.numpy().astype(np.int32)})[0]
# 2. Latent Initialization
latents = torch.randn([1, 4, 64, 64]).numpy()
# 3. UNet Denoising Loop (LCM 4-step)
print("Running UNet denoising...")
start_time = time.time()
for i, t in enumerate(timesteps):
noise_pred = unet_session.run(None, {
"sample": latents.astype(np.float32),
"/down_blocks.0/resnets.0/act_1/Mul_output_0": np.expand_dims(time_embeddings[i], axis=0),
"encoder_hidden_states": prompt_embeds
})[0]
# LCM Step Logic
alpha_prod_t = alphas_cumprod[t]
prev_t = timesteps[i + 1] if i < 3 else t
alpha_prod_t_prev = alphas_cumprod[prev_t] if i < 3 else final_alphas_cumprod
beta_prod_t = 1 - alpha_prod_t
# Boundary conditions
scaled_t = t * 10
c_skip = 0.5 ** 2 / (scaled_t ** 2 + 0.5 ** 2)
c_out = scaled_t / (scaled_t ** 2 + 0.5 ** 2) ** 0.5
pred_x0 = (latents - (beta_prod_t ** 0.5) * noise_pred) / (alpha_prod_t ** 0.5)
denoised = c_out * pred_x0 + c_skip * latents
if i < 3:
noise = torch.randn(noise_pred.shape).numpy()
latents = (alpha_prod_t_prev ** 0.5) * denoised + ((1 - alpha_prod_t_prev) ** 0.5) * noise
else:
latents = denoised
print(f"Denoising finished in {time.time() - start_time:.2f}s")
# 4. VAE Decoding
print("Decoding latents...")
latents = latents / 0.18215
image = vae_decoder.run(None, {"x": latents.astype(np.float32)})[0]
# 5. Post-processing & Save
image = np.transpose(image, (0, 2, 3, 1)).squeeze(0)
image = np.clip(image / 2 + 0.5, 0, 1)
image = (image * 255).astype("uint8")
pil_img = Image.fromarray(image)
pil_img.save(args.output)
print(f"Image saved to {args.output}")
if __name__ == "__main__":
main()