zjuJish commited on
Commit
a68e320
·
verified ·
1 Parent(s): ba4026b

Upload layer_diff_dataset/test_inp_4_try_index.py with huggingface_hub

Browse files
layer_diff_dataset/test_inp_4_try_index.py ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import torch
3
+ import os
4
+ from tqdm import tqdm
5
+ from modelscope.outputs import OutputKeys
6
+ from modelscope.pipelines import pipeline
7
+ from modelscope.utils.constant import Tasks
8
+ import argparse
9
+
10
+ # 创建命令行参数解析器
11
+ def parse_arguments():
12
+ parser = argparse.ArgumentParser(description="Image Inpainting with Pipeline")
13
+ parser.add_argument('--index', type=int, default=0,
14
+ help='Index for selecting images to process (default: 0)')
15
+ return parser.parse_args()
16
+
17
+ def main(args):
18
+ INDEX = args.index
19
+
20
+ # 进行图片 inpainting
21
+ prompt = 'hazy background with nothing on'
22
+
23
+ root_folder = '../data/P3M-10k/validation/P3M-500-NP'
24
+ jpeg_folder = os.path.join(root_folder, 'original_image')
25
+ mask_folder = os.path.join(root_folder, 'mask_dilate')
26
+ inp_folder = os.path.join(root_folder, 'inpainting')
27
+ os.makedirs(inp_folder, exist_ok=True)
28
+ vid_list = os.listdir(jpeg_folder)
29
+
30
+ image_inpainting = pipeline(
31
+ Tasks.image_inpainting,
32
+ model='/mnt/workspace/workgroup/sihui.jsh/alpha_work/diffusers/iic/cv_stable-diffusion-v2_image-inpainting_base',
33
+ device=f'cuda:{INDEX}', # 使用 index 来选择 GPU
34
+ torch_dtype=torch.float32,
35
+ enable_attention_slicing=True
36
+ )
37
+
38
+ pbar = tqdm(enumerate(vid_list), total=len(vid_list))
39
+ for i, vid_name in pbar:
40
+ if i < INDEX * 125:
41
+ continue
42
+ elif i >= INDEX * 125 + 125:
43
+ break
44
+ folder_path_0 = os.path.join(jpeg_folder, vid_name)
45
+ folder_path = os.path.join(mask_folder, vid_name.replace('.jpg', '.png'))
46
+ folder_path_ = os.path.join(inp_folder, vid_name)
47
+
48
+ if os.path.exists(folder_path_):
49
+ continue
50
+ h, w = cv2.imread(folder_path_0).shape[:2]
51
+ if h*w > 2000000:
52
+ image_tmp = cv2.resize(cv2.imread(folder_path_0), (int(w*0.8), int(h*0.8)))
53
+ cv2.imwrite('tmp.jpg', image_tmp)
54
+ folder_path_0 = 'tmp.jpg'
55
+ input_data = {
56
+ 'image': folder_path_0,
57
+ 'mask': folder_path,
58
+ 'prompt': prompt
59
+ }
60
+ output = image_inpainting(input_data)[OutputKeys.OUTPUT_IMG]
61
+ cv2.imwrite(folder_path_, output)
62
+
63
+ if __name__ == "__main__":
64
+ args = parse_arguments()
65
+ main(args)