Download code/jl_synth.py from fnruha0921/knps-change-detection-tmp: direct link, hf CLI and curl.
- Browser
- Download file 1.97 kB
-
https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/jl_synth.py
- Command line
-
hf download hf://fnruha0921/knps-change-detection-tmp/code/jl_synth.py
-
curl -L -o jl_synth.py https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/jl_synth.py
1.97 kB
| """Synth5: real JL1-CD pairs (pseudo class labels) with prob P_REAL, else Synth3 synthetic pairs.""" | |
| import os, time | |
| import cv2, numpy as np, torch | |
| from train2 import CROP | |
| from train3 import Synth3 | |
| class Synth5(Synth3): | |
| def __init__(self, n, p_real=0.5): | |
| super().__init__(n) | |
| self.p_real = p_real | |
| self.jpre = np.load("/workspace/prep/jl_pre.npy", mmap_mode="r"); self.jpost = np.load("/workspace/prep/jl_post.npy", mmap_mode="r") | |
| self.jb = np.load("/workspace/prep/jl_lb.npy", mmap_mode="r"); self.jt = np.load("/workspace/prep/jl_lt.npy", mmap_mode="r") | |
| print("JL1 real pairs", self.jpre.shape, flush=True) | |
| def __getitem__(self, idx): | |
| rng = np.random.default_rng((idx * 7919 + os.getpid() * 104729 + time.time_ns()) % 2**32) | |
| if rng.random() >= self.p_real: | |
| return super().__getitem__(idx) | |
| i = rng.integers(len(self.jpre)); size = int(CROP * rng.uniform(0.85, 1.15)); r, c = rng.integers(0, 512 - size + 1, 2) | |
| cr = lambda a, it: cv2.resize(np.ascontiguousarray(a[i, r:r + size, c:c + size]), (CROP, CROP), interpolation=it) | |
| pre, post = cr(self.jpre, cv2.INTER_AREA), cr(self.jpost, cv2.INTER_AREA) | |
| lb, lt = cr(self.jb, cv2.INTER_NEAREST), cr(self.jt, cv2.INTER_NEAREST) | |
| if rng.random() < 0.3: pre = self.photo(pre, rng) | |
| if rng.random() < 0.3: post = self.photo(post, rng) | |
| k = rng.integers(8) | |
| def d4(a): | |
| a = np.rot90(a, k % 4) | |
| return np.ascontiguousarray(a[:, ::-1] if k >= 4 else a) | |
| pre, post, lb, lt = d4(pre), d4(post), d4(lb), d4(lt) | |
| ig = np.full((CROP, CROP), 255, np.int64) | |
| t = lambda a: torch.from_numpy(a.transpose(2, 0, 1).copy()) | |
| lab = np.stack([lb, lt]).astype(np.float32); pres = np.array([lb.sum() >= 20, lt.sum() >= 20], np.float32) | |
| return (t(pre), t(post), torch.from_numpy(lab), torch.from_numpy(pres), torch.from_numpy(ig), torch.from_numpy(ig.copy())) | |