hideosnes/Zero-Shot-Material-Transfer
2
1from diffusers import StableDiffusionXLControlNetInpaintPipeline, ControlNetModel2from rembg import remove3from PIL import Image4import torch5from ip_adapter import IPAdapterXL6from ip_adapter.utils import register_cross_attention_hook, get_net_attn_map, attnmaps2images7from PIL import Image, ImageChops8from PIL import ImageEnhance9import numpy as np10import glob11 12def image_grid(imgs, rows, cols):13 assert len(imgs) == rows*cols14 15 w, h = imgs[0].size16 grid = Image.new('RGB', size=(cols*w, rows*h))17 grid_w, grid_h = grid.size18 19 for i, img in enumerate(imgs):20 grid.paste(img, box=(i%cols*w, i//cols*h))21 return grid22 23base_model_path = "stabilityai/stable-diffusion-xl-base-1.0"24image_encoder_path = "models/image_encoder"25ip_ckpt = "sdxl_models/ip-adapter_sdxl_vit-h.bin"26controlnet_path = "diffusers/controlnet-depth-sdxl-1.0"27device = "cuda"28 29torch.cuda.empty_cache()30 31# load SDXL pipeline32controlnet = ControlNetModel.from_pretrained(controlnet_path, variant="fp16", use_safetensors=True, torch_dtype=torch.float16).to(device)33pipe = StableDiffusionXLControlNetInpaintPipeline.from_pretrained(34 base_model_path,35 controlnet=controlnet,36 use_safetensors=True,37 torch_dtype=torch.float16,38 add_watermarker=False,39).to(device)40pipe.unet = register_cross_attention_hook(pipe.unet)41 42ip_model = IPAdapterXL(pipe, image_encoder_path, ip_ckpt, device)43 44 45 46textures = [tex.split('/')[-1].replace('.png', '') for tex in glob.glob('demo_assets/material_exemplars/*.png')]47objs = [obj.split('/')[-1].replace('.png', '') for obj in glob.glob('demo_assets/input_imgs/*.png')]48 49for texture in textures:50 for obj in objs:51 target_image_path = 'demo_assets/input_imgs/' + obj + '.png' # Replace with your image path52 target_image = Image.open(target_image_path).convert('RGB')53 rm_bg = remove(target_image)54 # output.save(output_path)55 target_mask = rm_bg.convert("RGB").point(lambda x: 0 if x < 1 else 255).convert('L').convert('RGB')# Convert mask to grayscale56 57 # Ensure mask is the same size as image58 59 # mask = ImageChops.invert(mask)60 # Generate random noise for the size of the image61 noise = np.random.randint(0, 256, target_image.size + (3,), dtype=np.uint8)62 noise_image = Image.fromarray(noise)63 mask_target_img = ImageChops.lighter(target_image, target_mask)64 invert_target_mask = ImageChops.invert(target_mask)65 66 67 gray_target_image = target_image.convert('L').convert('RGB')68 gray_target_image = ImageEnhance.Brightness(gray_target_image)69 70 # Adjust brightness71 # The factor 1.0 means original brightness, greater than 1.0 makes the image brighter. Adjust this if the image is too dim72 factor = 1.0 # Try adjusting this to get the desired brightness73 74 gray_target_image = gray_target_image.enhance(factor)75 grayscale_img = ImageChops.darker(gray_target_image, target_mask)76 img_black_mask = ImageChops.darker(target_image, invert_target_mask)77 grayscale_init_img = ImageChops.lighter(img_black_mask, grayscale_img)78 init_img = grayscale_init_img79 80 ip_image = Image.open("demo_assets/material_exemplars/" + texture + ".png")81 np_image = np.array(Image.open('demo_assets/depths/' + obj + '.png'))82 83 np_image = (np_image / 256).astype('uint8')84 85 depth_map = Image.fromarray(np_image).resize((1024,1024))86 87 init_img = init_img.resize((1024,1024))88 mask = target_mask.resize((1024, 1024))89 90 num_samples = 191 images = ip_model.generate(pil_image=ip_image, image=init_img, control_image=depth_map, mask_image=mask, controlnet_conditioning_scale=0.9, num_samples=num_samples, num_inference_steps=30, seed=42)92 images[0].save('demo_assets/output_images/' + obj + '_' + texture + '.png' )93 94 