File size: 4,958 Bytes
a3bbf71
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
"""Synth6 = Synth4 (EarthView negatives, p_ev) + HOT building scenes (p_hot of building scenes) + Korean-style vinyl greenhouse synthesis (p_gh)."""
import os, time, math
import cv2, numpy as np, torch
from train2 import CROP
from train3 import degrade
from ev_synth import Synth4

class Synth6(Synth4):
    def __init__(self, n, p_ev=0.2, p_hot=0.5, p_gh=0.12):
        super().__init__(n, p_ev)
        self.p_hot, self.p_gh = p_hot, p_gh
        self.hx = np.load("/workspace/prep/hot_x.npy", mmap_mode="r"); self.hy = np.load("/workspace/prep/hot_y.npy", mmap_mode="r")
        print("HOT scenes", self.hx.shape, flush=True)

    def scene(self, rng, kind):
        if kind == "b" and rng.random() < self.p_hot:
            i = rng.integers(len(self.hx)); s = rng.uniform(0.85, 1.15); size = min(256, int(round(CROP / s)))
            r, c = rng.integers(0, 257 - size, 2)
            x = cv2.resize(np.ascontiguousarray(self.hx[i, r:r + size, c:c + size]), (CROP, CROP), interpolation=cv2.INTER_AREA)
            y = cv2.resize(np.ascontiguousarray(self.hy[i, r:r + size, c:c + size]), (CROP, CROP), interpolation=cv2.INTER_NEAREST)
            return x.copy(), y.astype(np.uint8)          # code 1 = building (others unlabeled)
        return super().scene(rng, kind)

    @staticmethod
    def greenhouses(rng, shape):
        """Rows of long, narrow tunnel greenhouses: returns (mask, rendered rgb layer)."""
        m = np.zeros(shape, np.uint8); lay = np.zeros(shape + (3,), np.float32)
        th = rng.uniform(0, np.pi); u = np.array([np.cos(th), np.sin(th)]); v = np.array([-u[1], u[0]])
        n = int(rng.integers(2, 9)); w = rng.uniform(8, 14); gap = rng.uniform(1.5, 5); L = rng.uniform(30, 140)
        c0 = np.array([rng.uniform(40, shape[1] - 40), rng.uniform(40, shape[0] - 40)])
        base = rng.uniform(185, 245); tint = rng.uniform(-12, 12, 3)
        yy, xx = np.mgrid[:shape[0], :shape[1]].astype(np.float32)
        for k in range(n):
            ck = c0 + v * (k - n / 2) * (w + gap) + u * rng.uniform(-6, 6)
            Lk = L * rng.uniform(0.8, 1.0)
            pts = np.array([ck + u * Lk / 2 + v * w / 2, ck + u * Lk / 2 - v * w / 2, ck - u * Lk / 2 - v * w / 2, ck - u * Lk / 2 + v * w / 2], np.int32)
            mk = np.zeros(shape, np.uint8); cv2.fillPoly(mk, [pts], 1)
            dv = ((xx - ck[0]) * v[0] + (yy - ck[1]) * v[1]) / (w / 2)          # -1..1 across the tunnel
            shade = 0.78 + 0.22 * np.cos(np.clip(dv, -1, 1) * np.pi / 2)        # bright ridge, darker sides
            for ch in range(3):
                lay[..., ch] = np.where(mk > 0, (base + tint[ch]) * shade + rng.normal(0, 3, shape), lay[..., ch])
            m |= mk
        return m, np.clip(lay, 0, 255)

    def __getitem__(self, idx):
        rng = np.random.default_rng((idx * 7919 + os.getpid() * 104729 + time.time_ns()) % 2**32)
        if rng.random() >= self.p_gh:
            return super().__getitem__(idx)
        x, y = self.scene(rng, "g" if rng.random() < 0.6 else "any")
        pre = x.copy(); sem = np.zeros(y.shape, np.uint8); sem[y == 1] = 1; sem[y == 2] = 2
        if rng.random() < 0.3:   # pre-existing greenhouses in both images (not a change)
            m0, l0 = self.greenhouses(rng, y.shape); a = (m0[..., None] * rng.uniform(0.85, 0.95))
            pre = (pre * (1 - a) + l0 * a).astype(np.uint8); sem[m0 > 0] = 1
        post = pre.copy(); m, lay = self.greenhouses(rng, y.shape)
        m[sem == 1] = 0
        a = cv2.GaussianBlur(m.astype(np.float32), (0, 0), 0.6)[..., None] * rng.uniform(0.85, 0.97)
        post = (post * (1 - a) + lay * a).astype(np.uint8)
        lb = m.copy(); lt = np.zeros_like(m); sem_pre = sem.copy(); sem_post = sem.copy(); sem_post[m > 0] = 1
        if rng.random() < 0.15:   # demolition: reversed pair is not a target
            pre, post = post, pre; sem_pre, sem_post = sem_post, sem_pre; lb[:] = 0
        pre = degrade(self.photo(pre, rng), rng); post = degrade(self.photo(post, rng), rng)
        if rng.random() < 0.8:
            M = cv2.getRotationMatrix2D((CROP / 2, CROP / 2), rng.uniform(-1.5, 1.5), rng.uniform(0.98, 1.02)); M[:, 2] += rng.uniform(-4, 4, 2)
            pre = cv2.warpAffine(pre, M, (CROP, CROP), borderMode=cv2.BORDER_REFLECT)
            sem_pre = cv2.warpAffine(sem_pre, M, (CROP, CROP), flags=cv2.INTER_NEAREST, borderMode=cv2.BORDER_REFLECT)
        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, sem_pre, sem_post = map(d4, (pre, post, lb, lt, sem_pre, sem_post))
        t = lambda a: torch.from_numpy(a.transpose(2, 0, 1).copy())
        pres = np.array([lb.sum() >= 20, 0], np.float32)
        return (t(pre), t(post), torch.from_numpy(np.stack([lb, lt]).astype(np.float32)), torch.from_numpy(pres),
                torch.from_numpy(sem_pre.astype(np.int64)), torch.from_numpy(sem_post.astype(np.int64)))