zjuJish commited on
Commit
694dece
·
verified ·
1 Parent(s): d7b5444

Upload layer_diff_dataset/test_inp_4 copy 8.py with huggingface_hub

Browse files
layer_diff_dataset/test_inp_4 copy 8.py ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import torch
3
+ import os
4
+ import json
5
+ from tqdm import tqdm
6
+ from modelscope.outputs import OutputKeys
7
+ from modelscope.pipelines import pipeline
8
+ from modelscope.utils.constant import Tasks
9
+
10
+ # input_location = 'https://modelscope.oss-cn-beijing.aliyuncs.com/test/images/image_inpainting/image_inpainting_1.png'
11
+ # input_mask_location = 'https://modelscope.oss-cn-beijing.aliyuncs.com/test/images/image_inpainting/image_inpainting_mask_1.png'
12
+ prompt = 'hazy background with nothing on'
13
+
14
+ root_folder = '../data/video_dataset/YoutubeVOS/train'
15
+ jpeg_folder = os.path.join(root_folder,'JPEGImages')
16
+ mask_folder = os.path.join(root_folder,'mask_dilate')
17
+ inp_folder = os.path.join(root_folder,'inp_image_256')
18
+ os.makedirs(inp_folder,exist_ok=True)
19
+ vid_list = os.listdir(jpeg_folder)
20
+
21
+ image_inpainting = pipeline(
22
+ Tasks.image_inpainting,
23
+ model='/mnt/workspace/workgroup/sihui.jsh/alpha_work/diffusers/iic/cv_stable-diffusion-v2_image-inpainting_base',
24
+ device='cuda:4',
25
+ torch_dtype=torch.float32,
26
+ enable_attention_slicing=True)
27
+
28
+ pbar = tqdm(enumerate(vid_list),total=len(vid_list))
29
+ for i, vid_name in pbar:
30
+ if i<=1600:
31
+ continue
32
+ if i>1800:
33
+ break
34
+ # folder_path_0 = 'YoutubeVOS/JPEGImages/0043f083b5'
35
+ folder_path_0 = os.path.join(jpeg_folder,vid_name)
36
+ folder_path = os.path.join(mask_folder,vid_name)
37
+ # folder_path = 'YoutubeVOS/mask_dilate/0043f083b5'
38
+ # folder_path_ = 'YoutubeVOS/inp/0043f083b5'
39
+ folder_path_ = os.path.join(inp_folder,vid_name)
40
+ os.makedirs(folder_path_,exist_ok=True)
41
+ file_list = os.listdir(folder_path)
42
+ file_list = [i for i in file_list if i.endswith('.png')]
43
+ file_list.sort()
44
+
45
+ # pbar = tqdm(enumerate(file_list),total=len(file_list))
46
+ # with open('/mnt/workspace/workgroup/sihui.jsh/layer_diff_dataset/train/im_rgba.json', 'r') as file:
47
+ # data = json.load(file)
48
+
49
+ for i, image_name in enumerate(file_list):
50
+ if os.path.exists(os.path.join(folder_path_,image_name)):
51
+ continue
52
+ last_idx = min(15,len(file_list)-1)
53
+ if not (i == 0 or i == last_idx):
54
+ continue
55
+ # print('i',i)
56
+ # if i<1000:
57
+ # continue
58
+ # if i>500:
59
+ # break
60
+
61
+ # if i==1:
62
+ # break
63
+ # gt_name = data[i]["images"].split('/')[-1]
64
+ # print(image_name,gt_name)
65
+ # prompt = data[i]["prompt_fg"]
66
+ # print(prompt)
67
+
68
+ # mask_image = load_image(os.path.join(folder_path,image_name)).resize((512, 512))
69
+ # image = load_image(os.path.join(folder_path_0,image_name.split('.png')[0]+'.jpg')).resize((512, 512))
70
+
71
+ input = {
72
+ 'image': os.path.join(folder_path_0,image_name.split('.png')[0]+'.jpg'),
73
+ 'mask': os.path.join(folder_path,image_name),
74
+ 'prompt': prompt
75
+ }
76
+ output = image_inpainting(input)[OutputKeys.OUTPUT_IMG]
77
+ output = cv2.resize(output,(256,256))
78
+ cv2.imwrite(os.path.join(folder_path_,image_name), output)
79
+ # exit(0)