fnruha0921's picture
add code
eea5f0e verified
Raw History Blame Contribute Delete
10.3 kB
"""Build submission notebooks (unet = baseline arch, siam = SiamCD) from the baseline notebook."""
import json
import sys
BASE = "baseline/03_illegal structure submission/predict.ipynb"
kind, out, th_b, th_t, area_b, area_t, pres_t = sys.argv[1], sys.argv[2], *map(float, sys.argv[3:8])
ens_w = float(sys.argv[8]) if len(sys.argv) > 8 else 0.5 # weight of the UNet in the ensemble
nb = json.load(open(BASE))
cells = [c for c in nb["cells"] if "aifactory submit" not in "".join(c["source"])]
for c in cells:
if c["cell_type"] == "code":
c["outputs"], c["execution_count"] = [], None
setup = "".join(cells[1]["source"])
setup = setup.replace('CKPT = Path("assets/model/unet_r18_cd.pt")',
'CKPT = next(p for p in map(Path, ["assets/model/unet_synth.pt", "assets/model/siam_v2.pt", "assets/model/hybrid_v3.pt"]) if p.exists())')
setup = setup.replace("MIN_AREA = 30 ", f"MIN_AREA = {{'new_building': {area_b}, 'tree_removal': {area_t}}} ")
setup = setup.replace("def mask_to_polygons(mask: np.ndarray) -> str:", "def mask_to_polygons(mask: np.ndarray, min_area: float) -> str:")
setup = setup.replace("p.area >= MIN_AREA]", "p.area >= min_area]")
setup += f"""
THRESH = {{'new_building': {th_b}, 'tree_removal': {th_t}}} # pixel probability thresholds
PRES_THRESH = {pres_t} # presence-head gate (siam only, 0 = off)
"""
cells[1]["source"] = setup
if kind == "unet":
model_src = '''# [model]
import segmentation_models_pytorch as smp
def load_model(device: str):
model = smp.Unet(encoder_name="resnet18", encoder_weights=None, in_channels=6, classes=3)
ck = torch.load(CKPT, map_location="cpu", weights_only=True)
model.load_state_dict(ck["state_dict"])
return model.to(device).eval()
def read_image(path: Path) -> np.ndarray:
with Image.open(io.BytesIO(path.read_bytes())) as im:
arr = np.asarray(im.convert("RGB"))
if arr.shape[:2] != (H, W):
raise ValueError(f"{path}: 크기 {arr.shape[:2]} κ°€ {(H, W)} 와 λ‹€λ¦…λ‹ˆλ‹€")
return arr
@torch.no_grad()
def predict_batch(model, pres, posts, device):
"""-> change prob (N,2,H,W) for new_building/tree_removal, presence prob (N,2) or None. 4-way flip TTA."""
x = np.stack([np.concatenate([(a / 255.0 - MEAN) / STD, (b / 255.0 - MEAN) / STD], axis=2)
for a, b in zip(pres, posts)]).astype(np.float32)
x = torch.from_numpy(x).permute(0, 3, 1, 2).to(device)
prob = 0
for dims in ([], [3], [2], [2, 3]):
xi = torch.flip(x, dims) if dims else x
pi = torch.softmax(model(xi), dim=1)[:, 1:3]
prob = prob + (torch.flip(pi, dims) if dims else pi) / 4
return prob.float().cpu().numpy(), None
'''
elif kind == "hybrid":
model_src = '''# [model]
sys.path.insert(0, "assets")
from model_v3 import HybridCD
def load_model(device: str):
model = HybridCD()
ck = torch.load(CKPT, map_location="cpu", weights_only=True)
model.load_state_dict(ck["state_dict"])
return model.to(device).eval()
def read_image(path: Path) -> np.ndarray:
with Image.open(io.BytesIO(path.read_bytes())) as im:
arr = np.asarray(im.convert("RGB"))
if arr.shape[:2] != (H, W):
raise ValueError(f"{path}: 크기 {arr.shape[:2]} κ°€ {(H, W)} 와 λ‹€λ¦…λ‹ˆλ‹€")
return arr
@torch.no_grad()
def predict_batch(model, pres, posts, device):
a = torch.from_numpy(np.stack(pres)).permute(0, 3, 1, 2).to(device).float() / 255
b = torch.from_numpy(np.stack(posts)).permute(0, 3, 1, 2).to(device).float() / 255
prob, pp = 0, 0
for dims in ([], [3], [2], [2, 3]):
ai, bi = (torch.flip(a, dims), torch.flip(b, dims)) if dims else (a, b)
with torch.autocast("cuda", dtype=torch.float16, enabled=(device == "cuda")):
ch, pr, _, _, _ = model(ai, bi)
ci = ch.float().sigmoid()
prob = prob + (torch.flip(ci, dims) if dims else ci) / 4
pp = pp + pr.float().sigmoid() / 4
return prob.cpu().numpy(), pp.cpu().numpy()
'''
elif kind == "ens":
model_src = f'''# [model]
sys.path.insert(0, "assets")
import segmentation_models_pytorch as smp
from model_v2 import SiamCD
ENS_W = {ens_w} # weight of the baseline-arch UNet; 1-ENS_W for the Siamese model
def load_model(device: str):
u = smp.Unet(encoder_name="resnet18", encoder_weights=None, in_channels=6, classes=3)
u.load_state_dict(torch.load("assets/model/unet_synth.pt", map_location="cpu", weights_only=True)["state_dict"])
s = SiamCD()
s.load_state_dict(torch.load("assets/model/siam_v2.pt", map_location="cpu", weights_only=True)["state_dict"])
return (u.to(device).eval(), s.to(device).eval())
def read_image(path: Path) -> np.ndarray:
with Image.open(io.BytesIO(path.read_bytes())) as im:
arr = np.asarray(im.convert("RGB"))
if arr.shape[:2] != (H, W):
raise ValueError(f"{{path}}: 크기 {{arr.shape[:2]}} κ°€ {{(H, W)}} 와 λ‹€λ¦…λ‹ˆλ‹€")
return arr
@torch.no_grad()
def predict_batch(model, pres, posts, device):
u, s = model
a = torch.from_numpy(np.stack(pres)).permute(0, 3, 1, 2).to(device).float() / 255
b = torch.from_numpy(np.stack(posts)).permute(0, 3, 1, 2).to(device).float() / 255
m = torch.tensor(MEAN, device=device).view(1, 3, 1, 1); sd = torch.tensor(STD, device=device).view(1, 3, 1, 1)
prob, pp = 0, 0
for dims in ([], [3], [2], [2, 3]):
ai, bi = (torch.flip(a, dims), torch.flip(b, dims)) if dims else (a, b)
pu = torch.softmax(u(torch.cat([(ai - m) / sd, (bi - m) / sd], 1)), 1)[:, 1:3]
with torch.autocast("cuda", dtype=torch.float16, enabled=(device == "cuda")):
ch, pr, _, _ = s(ai, bi)
ci = ENS_W * pu.float() + (1 - ENS_W) * ch.float().sigmoid()
prob = prob + (torch.flip(ci, dims) if dims else ci) / 4
pp = pp + pr.float().sigmoid() / 4
return prob.cpu().numpy(), pp.cpu().numpy()
'''
else:
model_src = '''# [model]
sys.path.insert(0, "assets")
from model_v2 import SiamCD
def load_model(device: str):
model = SiamCD()
ck = torch.load(CKPT, map_location="cpu", weights_only=True)
model.load_state_dict(ck["state_dict"])
return model.to(device).eval()
def read_image(path: Path) -> np.ndarray:
with Image.open(io.BytesIO(path.read_bytes())) as im:
arr = np.asarray(im.convert("RGB"))
if arr.shape[:2] != (H, W):
raise ValueError(f"{path}: 크기 {arr.shape[:2]} κ°€ {(H, W)} 와 λ‹€λ¦…λ‹ˆλ‹€")
return arr
@torch.no_grad()
def predict_batch(model, pres, posts, device):
"""-> change prob (N,2,H,W), presence prob (N,2). 4-way flip TTA."""
a = torch.from_numpy(np.stack(pres)).permute(0, 3, 1, 2).to(device).float() / 255
b = torch.from_numpy(np.stack(posts)).permute(0, 3, 1, 2).to(device).float() / 255
prob, pp = 0, 0
for dims in ([], [3], [2], [2, 3]):
ai, bi = (torch.flip(a, dims), torch.flip(b, dims)) if dims else (a, b)
with torch.autocast("cuda", dtype=torch.float16, enabled=(device == "cuda")):
ch, pr, _, _ = model(ai, bi)
ci = ch.float().sigmoid()
prob = prob + (torch.flip(ci, dims) if dims else ci) / 4
pp = pp + pr.float().sigmoid() / 4
return prob.cpu().numpy(), pp.cpu().numpy()
'''
model_src += '''
DEVICE = "cuda" if torch.cuda.is_available() else ("mps" if torch.backends.mps.is_available() else "cpu")
try:
MODEL = load_model(DEVICE)
predict_batch(MODEL, [np.zeros((H, W, 3), np.uint8)], [np.zeros((H, W, 3), np.uint8)], DEVICE)
except Exception as e: # noqa: BLE001
print(f"{DEVICE} μ‹œν—˜ 좔둠에 μ‹€νŒ¨ν•΄({type(e).__name__}: {e}) CPU 둜 μ „ν™˜ν•©λ‹ˆλ‹€")
DEVICE = "cpu"
MODEL = load_model(DEVICE)
print("device", DEVICE)
'''
cells[4]["source"] = model_src
cells[5]["source"] = '''# [infer]
ROWS = []
STATS = {c: {"maxp": [], "pres": []} for c in CLASSES}
GRID_T = (0.3, 0.4, 0.5, 0.6, 0.7, 0.8)
GRID_A = (20, 50, 100, 200)
GRID = {c: np.zeros((len(GRID_T), len(GRID_A)), int) for c in CLASSES}
t0 = time.time()
for s in range(0, len(IDS), BATCH):
ids = IDS[s:s + BATCH]
pres = [read_image(ROOT / "images" / i / "pre.png") for i in ids]
posts = [read_image(ROOT / "images" / i / "post.png") for i in ids]
prob, pp = predict_batch(MODEL, pres, posts, DEVICE)
nd = np.stack([(a.max(2) == 0) | (b.max(2) == 0) for a, b in zip(pres, posts)]) # no-data is never scored
prob[:, :, :, :][np.repeat(nd[:, None], 2, 1)] = 0
for j, i in enumerate(ids):
cells = []
for k, c in enumerate(CLASSES):
p = prob[j, k]
STATS[c]["maxp"].append(float(p.max()))
if pp is not None:
STATS[c]["pres"].append(float(pp[j, k]))
for ti, t in enumerate(GRID_T):
area = int((p > t).sum())
for ai, amin in enumerate(GRID_A):
GRID[c][ti, ai] += area >= amin
gate = pp is None or PRES_THRESH <= 0 or pp[j, k] >= PRES_THRESH
cells.append(mask_to_polygons(p > THRESH[c], MIN_AREA[c]) if gate else "")
ROWS.append((i, *cells))
if (s // BATCH) % 10 == 0 or s + BATCH >= len(IDS):
print(f" {min(s + BATCH, len(IDS))}/{len(IDS)} {time.time() - t0:.1f}s", flush=True)
n_b = sum(1 for _, b, _ in ROWS if b)
n_t = sum(1 for _, _, t in ROWS if t)
print(f"좔둠을 μ™„λ£Œν–ˆμŠ΅λ‹ˆλ‹€: {len(ROWS)}건, 증좕 μ–‘μ„± {n_b}건, 벌λͺ© μ–‘μ„± {n_t}건, {time.time() - t0:.1f}s")
print("μ„€μ •:", "THRESH", THRESH, "MIN_AREA", MIN_AREA, "PRES_THRESH", PRES_THRESH)
for c in CLASSES:
mp = np.array(STATS[c]["maxp"])
print(f"[{c}] μŒλ³„ μ΅œλŒ€ν™•λ₯  λΆ„μœ„μˆ˜(10/25/50/75/90%):", np.round(np.percentile(mp, [10, 25, 50, 75, 90]), 3).tolist())
if STATS[c]["pres"]:
pr = np.array(STATS[c]["pres"])
print(f"[{c}] μ‘΄μž¬ν™•λ₯  λΆ„μœ„μˆ˜(10/25/50/75/90%):", np.round(np.percentile(pr, [10, 25, 50, 75, 90]), 3).tolist(),
" >=0.3/0.5/0.7:", [int((pr >= v).sum()) for v in (0.3, 0.5, 0.7)])
print(f"[{c}] μ–‘μ„± 쌍 수 ν‘œ (ν–‰=ν”½μ…€μž„κ³„ {GRID_T}, μ—΄=μ΅œμ†Œλ©΄μ  {GRID_A}):")
for ti, t in enumerate(GRID_T):
print(" ", t, GRID[c][ti].tolist())
'''
nb["cells"] = cells
json.dump(nb, open(out, "w"), ensure_ascii=False, indent=1)
print("wrote", out)