randomwalkers commited on
Commit
3a48e4d
·
verified ·
1 Parent(s): 2b4669c

Upload folder using huggingface_hub

Browse files
Files changed (7) hide show
  1. blend.py +173 -0
  2. infer_ir.py +18 -4
  3. read_from_dir.py +1 -1
  4. readme.md +3 -0
  5. requirements.txt +2 -1
  6. run.sh +8 -3
  7. unet_178.pth +3 -0
blend.py ADDED
@@ -0,0 +1,173 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ import argparse
3
+ from pathlib import Path
4
+
5
+ import cv2
6
+ import numpy as np
7
+
8
+
9
+ def parse_args():
10
+ parser = argparse.ArgumentParser(
11
+ description=(
12
+ "Canny edge mix: edge regions use refine image, "
13
+ "non-edge regions use weighted average."
14
+ )
15
+ )
16
+ parser.add_argument("--base-dir", type=str, required=True, help="Directory of base images.")
17
+ parser.add_argument("--refine-dir", type=str, required=True, help="Directory of refine images.")
18
+ parser.add_argument("--out-dir", type=str, required=True, help="Output root directory.")
19
+ parser.add_argument(
20
+ "--exts",
21
+ type=str,
22
+ nargs="+",
23
+ default=["png", "jpg", "jpeg", "bmp", "webp"],
24
+ help="Image extensions to match (case-insensitive).",
25
+ )
26
+ parser.add_argument(
27
+ "--limit",
28
+ type=int,
29
+ default=0,
30
+ help="Only process first N matched pairs (0 means all).",
31
+ )
32
+
33
+ # Keep these five tunable parameters only.
34
+ parser.add_argument(
35
+ "--nonedge-alpha",
36
+ type=float,
37
+ default=0.5,
38
+ help="Base-image weight for non-edge averaging (0~1).",
39
+ )
40
+ parser.add_argument("--canny-low", type=int, default=80, help="Canny lower threshold.")
41
+ parser.add_argument("--canny-high", type=int, default=160, help="Canny upper threshold.")
42
+ parser.add_argument(
43
+ "--canny-dilate",
44
+ type=int,
45
+ default=1,
46
+ help="Dilate iterations for edge mask.",
47
+ )
48
+ parser.add_argument(
49
+ "--edge-feather",
50
+ type=int,
51
+ default=1,
52
+ help="Gaussian feather radius for edge mask (0 means hard edge).",
53
+ )
54
+ return parser.parse_args()
55
+
56
+
57
+ def build_pairs(base_dir: Path, refine_dir: Path, exts):
58
+ exts_set = {e.lower().lstrip(".") for e in exts}
59
+ base_map = {}
60
+ for p in base_dir.rglob("*"):
61
+ if p.is_file() and p.suffix.lower().lstrip(".") in exts_set:
62
+ rel = p.relative_to(base_dir).as_posix()
63
+ base_map[rel] = p
64
+
65
+ pairs = []
66
+ missing = 0
67
+ for rel, base_path in sorted(base_map.items()):
68
+ refine_path = refine_dir / rel
69
+ if refine_path.exists():
70
+ pairs.append((rel, base_path, refine_path))
71
+ else:
72
+ missing += 1
73
+ return pairs, missing
74
+
75
+
76
+ def canny_edge_mask(base_bgr_u8, canny_low, canny_high, dilate_iter=1, feather_radius=1):
77
+ gray = cv2.cvtColor(base_bgr_u8, cv2.COLOR_BGR2GRAY)
78
+ gray = cv2.GaussianBlur(gray, (5, 5), 0.0)
79
+ edge = cv2.Canny(gray, int(canny_low), int(canny_high))
80
+ if dilate_iter > 0:
81
+ kernel = np.ones((3, 3), np.uint8)
82
+ edge = cv2.dilate(edge, kernel, iterations=int(dilate_iter))
83
+ edge01 = edge.astype(np.float32) / 255.0
84
+ if feather_radius > 0:
85
+ k = int(feather_radius) * 2 + 1
86
+ edge01 = cv2.GaussianBlur(edge01, (k, k), 0.0)
87
+ edge01 = np.clip(edge01, 0.0, 1.0)
88
+ return edge01, edge
89
+
90
+
91
+ def blend_with_base_weight(base_bgr_u8, refine_bgr_u8, base_weight_01):
92
+ base = base_bgr_u8.astype(np.float32) / 255.0
93
+ refine = refine_bgr_u8.astype(np.float32) / 255.0
94
+ w = base_weight_01[:, :, None]
95
+ out = w * base + (1.0 - w) * refine
96
+ out = np.clip(out * 255.0 + 0.5, 0.0, 255.0).astype(np.uint8)
97
+ return out
98
+
99
+
100
+ def ensure_parent(path: Path):
101
+ path.parent.mkdir(parents=True, exist_ok=True)
102
+
103
+
104
+ def main():
105
+ args = parse_args()
106
+ base_dir = Path(args.base_dir)
107
+ refine_dir = Path(args.refine_dir)
108
+ out_dir = Path(args.out_dir)
109
+
110
+ out_blend = out_dir / "blended"
111
+ out_mask = out_dir / "mask01"
112
+ out_mask_vis = out_dir / "mask_vis"
113
+ out_blend.mkdir(parents=True, exist_ok=True)
114
+ out_mask.mkdir(parents=True, exist_ok=True)
115
+ out_mask_vis.mkdir(parents=True, exist_ok=True)
116
+
117
+ pairs, missing = build_pairs(base_dir, refine_dir, args.exts)
118
+ if args.limit > 0:
119
+ pairs = pairs[: args.limit]
120
+ print(f"Found pairs: {len(pairs)}, missing in refine: {missing}")
121
+ if not pairs:
122
+ raise RuntimeError("No matching image pairs found.")
123
+
124
+ fail_count = 0
125
+ alpha_nonedge = float(np.clip(args.nonedge_alpha, 0.0, 1.0))
126
+ for idx, (rel, base_path, refine_path) in enumerate(pairs, 1):
127
+ base = cv2.imread(str(base_path), cv2.IMREAD_COLOR)
128
+ refine = cv2.imread(str(refine_path), cv2.IMREAD_COLOR)
129
+ if base is None or refine is None:
130
+ print(f"[WARN] Failed to read: {rel}")
131
+ fail_count += 1
132
+ continue
133
+ if base.shape != refine.shape:
134
+ print(f"[WARN] Shape mismatch, skip: {rel}, {base.shape} vs {refine.shape}")
135
+ fail_count += 1
136
+ continue
137
+
138
+ edge01, edge_bin = canny_edge_mask(
139
+ base,
140
+ args.canny_low,
141
+ args.canny_high,
142
+ dilate_iter=args.canny_dilate,
143
+ feather_radius=args.edge_feather,
144
+ )
145
+
146
+ # Edge: use refine directly (base weight = 0).
147
+ # Non-edge: weighted average by nonedge-alpha.
148
+ base_weight = np.clip((1.0 - edge01) * alpha_nonedge, 0.0, 1.0)
149
+ out = blend_with_base_weight(base, refine, base_weight)
150
+
151
+ rel_path = Path(rel)
152
+ p_blend = out_blend / rel_path
153
+ p_mask = out_mask / rel_path
154
+ p_mvis = out_mask_vis / rel_path
155
+ ensure_parent(p_blend)
156
+ ensure_parent(p_mask)
157
+ ensure_parent(p_mvis)
158
+
159
+ cv2.imwrite(str(p_blend), out)
160
+ cv2.imwrite(str(p_mask), (base_weight * 255.0 + 0.5).astype(np.uint8))
161
+
162
+ heat = cv2.applyColorMap(edge_bin.astype(np.uint8), cv2.COLORMAP_JET)
163
+ vis = cv2.addWeighted(refine, 0.78, heat, 0.22, 0.0)
164
+ cv2.imwrite(str(p_mvis), vis)
165
+
166
+ if idx % 100 == 0 or idx == len(pairs):
167
+ print(f"Processed {idx}/{len(pairs)}")
168
+
169
+ print(f"Done. failures={fail_count}, output={out_dir}")
170
+
171
+
172
+ if __name__ == "__main__":
173
+ main()
infer_ir.py CHANGED
@@ -17,7 +17,7 @@ weight_dtype = torch.bfloat16 # 可改为 torch.float16 / torch.float32
17
 
18
  vae_path = BASE_DIR / "vae.pth"
19
  transformer_path = BASE_DIR / "transformer.pth"
20
- unet_path = BASE_DIR / "unet.pth"
21
  text_embedding_path = BASE_DIR / "text_embeddings.pth"
22
 
23
  output_folder = BASE_DIR / "output"
@@ -111,6 +111,13 @@ if not input_folder.exists() or not input_folder.is_dir():
111
  raise ValueError(f"Invalid input folder: {input_folder}")
112
 
113
  output_folder.mkdir(parents=True, exist_ok=True)
 
 
 
 
 
 
 
114
  image_paths = collect_image_paths(input_folder)
115
 
116
  shift_factor = getattr(vae.config, "shift_factor", 0.0)
@@ -167,9 +174,16 @@ for img_path in tqdm(image_paths, desc="Processing"):
167
  x_pred_final = x_pred + unet_pred
168
 
169
  # ===== 保存 =====
170
- out_pil = vae_output_to_pil(x_pred_final)
171
  out_name = img_path.stem + ".png"
172
- out_path = output_folder / out_name
173
- out_pil.save(out_path, format="PNG", compress_level=0)
 
 
 
 
 
 
 
 
174
 
175
  print(f"Done. Saved to: {output_folder}")
 
17
 
18
  vae_path = BASE_DIR / "vae.pth"
19
  transformer_path = BASE_DIR / "transformer.pth"
20
+ unet_path = BASE_DIR / "unet_178.pth"
21
  text_embedding_path = BASE_DIR / "text_embeddings.pth"
22
 
23
  output_folder = BASE_DIR / "output"
 
111
  raise ValueError(f"Invalid input folder: {input_folder}")
112
 
113
  output_folder.mkdir(parents=True, exist_ok=True)
114
+
115
+ # 新增两个子文件夹
116
+ x_pred_folder = output_folder / "x_pred"
117
+ x_pred_final_folder = output_folder / "x_pred_final"
118
+ x_pred_folder.mkdir(parents=True, exist_ok=True)
119
+ x_pred_final_folder.mkdir(parents=True, exist_ok=True)
120
+
121
  image_paths = collect_image_paths(input_folder)
122
 
123
  shift_factor = getattr(vae.config, "shift_factor", 0.0)
 
174
  x_pred_final = x_pred + unet_pred
175
 
176
  # ===== 保存 =====
 
177
  out_name = img_path.stem + ".png"
178
+
179
+ # 保存 x_pred
180
+ x_pred_pil = vae_output_to_pil(x_pred)
181
+ x_pred_path = x_pred_folder / out_name
182
+ x_pred_pil.save(x_pred_path, format="PNG", compress_level=0)
183
+
184
+ # 保存 x_pred_final
185
+ x_pred_final_pil = vae_output_to_pil(x_pred_final)
186
+ x_pred_final_path = x_pred_final_folder / out_name
187
+ x_pred_final_pil.save(x_pred_final_path, format="PNG", compress_level=0)
188
 
189
  print(f"Done. Saved to: {output_folder}")
read_from_dir.py CHANGED
@@ -74,7 +74,7 @@ def process_images_from_dir(input_dir, output_base_dir=None):
74
 
75
  if __name__ == "__main__":
76
  # 配置变量
77
- input_directory = "./output" # 修改为你的输入目录路径
78
  output_base_directory = "./" # 如果为None,输出目录会在输入目录的父目录下创建
79
 
80
  process_images_from_dir(input_directory, output_base_directory)
 
74
 
75
  if __name__ == "__main__":
76
  # 配置变量
77
+ input_directory = "./final_result/blended" # 修改为你的输入目录路径
78
  output_base_directory = "./" # 如果为None,输出目录会在输入目录的父目录下创建
79
 
80
  process_images_from_dir(input_directory, output_base_directory)
readme.md CHANGED
@@ -8,3 +8,6 @@ pip install -r requirements.txt -i https://mirrors.cloud.tencent.com/pypi/simple
8
 
9
  ### 按照run.sh命令运行即可
10
  cd到当前项目路径,运行run.sh即可
 
 
 
 
8
 
9
  ### 按照run.sh命令运行即可
10
  cd到当前项目路径,运行run.sh即可
11
+ 最后的结果在processed_jpg目录下
12
+
13
+ 如果运行出现问题,可以联系:chenyx.cs@gmail.com
requirements.txt CHANGED
@@ -1,3 +1,4 @@
1
  diffusers==0.34.0
2
  torch==2.4.1
3
- torchvision==0.19.1
 
 
1
  diffusers==0.34.0
2
  torch==2.4.1
3
+ torchvision==0.19.1
4
+ opencv-python
run.sh CHANGED
@@ -1,7 +1,12 @@
1
  # input_folder 换成本地的输入路径,结果产生在./output
2
  python ./infer_ir.py \
3
- --input_folder /mnt/tidal-sh01/dataset/media_image_algo_tidalfs/zhongyilian/Datasets/Test/LoViF_LQ
4
 
5
- # 将png结果转为比赛要求的jpg
6
- python read_from_dir.py
 
 
 
7
 
 
 
 
1
  # input_folder 换成本地的输入路径,结果产生在./output
2
  python ./infer_ir.py \
3
+ --input_folder xxxxxxxxxxxxxxx
4
 
5
+ python ./blend.py \
6
+ --base-dir ./output/x_pred \
7
+ --refine-dir ./output/x_pred_final \
8
+ --out-dir ./final_result \
9
+ --nonedge-alpha 0.5 --canny-low 80 --canny-high 160 --canny-dilate 1 --edge-feather 1
10
 
11
+ # 将png结果转为比赛要求的jpg
12
+ python read_from_dir.py
unet_178.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3a70f6ea32f577af5573612956d69b3101444fae98575ba6835e10a56410b2e2
3
+ size 454870132