#!/usr/bin/env python3 # -*- coding: utf-8 -*- """ Region Growing (static R-tree, bbox+pixel list 저장) 개선 개요 - 씨드: grow_mask(느슨한 BTD 조건) 내부의 IR105 국소 최소에서 시작 -> 미성숙 구름도 씨드 생성됨 - 성장: 아래 제약으로 '큰 구름에 붙는 warm ridge'와 '가느다란 bridge'를 차단 1) tmin cap: tv <= tmin + 35 K 2) BTD2(=WV063-IR105) 유사성: |BTD2 - BTD2_seed| <= 8 K 3) 이웃 지지: 3x3 내 기존 클러스터 픽셀 >= 2 4) 연결성: 4-연결(권장) 5) 최소 픽셀 수 미달 시 롤백(폐기) * grow 마스크(느슨): BTD(105-123) < 8 K & BTD(063-105) > -50 K * seed 생성: grow_mask 내부에서 IR105 국소 최소 (radius=1) * 결과물 - region_id 2-D float32 (netCDF4, zlib 9) (클러스터 id, 배경 NaN) - *_clusters.pkl {cid:{bbox:[c1,r1,c2,r2], mask:bool2D, pixel_count:int, cappi:2D}} - *_rtree.{idx,dat} (bbox 인덱스) 경로 설정 - IN_ROOT/OUT_ROOT는 import 시점에는 빈 문자열이다. - main에서 build/run/step1_region_growing_config.json을 읽은 뒤 apply_config()가 실제 경로로 채운다. - bt_root가 data_preprocess/result/res_2km처럼 L1B/L2 폴더이면, 실제 labeling 입력인 L1B/YYYYMMDD를 자동 선택한다. """ from __future__ import annotations import os import sys import time import pickle import logging import argparse from datetime import datetime from typing import List, Tuple, Dict, Any, Optional from collections import deque import numpy as np import xarray as xr # 출력용 NetCDF 저장을 위해 xarray는 유지합니다. from rtree import index as rtree_index from scipy import ndimage as ndi # ───────────── 설정 및 로깅 ───────────── try: sys.stdout.reconfigure(encoding="utf-8", errors="replace") sys.stderr.reconfigure(encoding="utf-8", errors="replace") except AttributeError: pass CONFIG_NAME = "step1_region_growing_config.json" def _find_package_root(config_name: str) -> str: cur = os.path.dirname(os.path.abspath(__file__)) while True: if os.path.exists(os.path.join(cur, "build", "run", config_name)): return cur parent = os.path.dirname(cur) if parent == cur: return os.path.abspath(os.getcwd()) cur = parent ROOT = _find_package_root(CONFIG_NAME) RUN_DIR = os.path.join(ROOT, "build", "run") sys.path.insert(0, ROOT) from .config_utils import load_config def _resolve_path(path: str) -> str: if os.path.isabs(path): return path return os.path.abspath(os.path.join(RUN_DIR, path)) def _has_date_dirs(path: str) -> bool: if not os.path.isdir(path): return False return any( os.path.isdir(os.path.join(path, name)) and len(name) == 8 and name.isdigit() for name in os.listdir(path) ) def _resolve_data_root(path: str) -> str: if _has_date_dirs(path): return path for subdir in ("L1B", "l1b"): candidate = os.path.join(path, subdir) if _has_date_dirs(candidate): return candidate return path def _set_logging(log_file: str) -> None: log_dir = os.path.dirname(log_file) if log_dir: os.makedirs(log_dir, exist_ok=True) logger.handlers.clear() logger.setLevel(logging.INFO) fmt = logging.Formatter("%(asctime)s - %(levelname)s - %(message)s") for handler in (logging.FileHandler(log_file, encoding="utf-8"), logging.StreamHandler()): handler.setFormatter(fmt) logger.addHandler(handler) # Config 적용 전 placeholder. 실제 값은 apply_config()에서 채운다. IN_ROOT = "" OUT_ROOT = "" EXCLUDE_VAL_FOLDER = True START_DATE = "20250701" # 20200730 20210809 20220815 20230814 20240901 END_DATE = "20250731" DATE_PAIRS: List[Tuple[str, str]] = [] # 비우면 START_DATE~END_DATE 처리 logger = logging.getLogger("PERSIANN_CCS_SEG") PROFILE_TIME = True # ───────────── (1) cloud mask / ITT params ───────────── DT_K = 3.0 TU_K = 285.0 TU_ITT = 285.0 CONNECTIVITY = 8 # 4 or 8 MIN_PIXELS_KEEP = 2 USE_SIMPLE_CLOUD_MASK = True # Tb <= TU_K USE_BTD_MASK = True BTD2_GROW = -47.77727170098097 # (wv063 - ir105) > BTD2_GROW # ───────────── seed params (ITT Step1) ───────────── SEED_WIN = 9 SEED_MERGE_PIX = 2 SEED_MIN_PIX = 2 def apply_config(cfg: Dict[str, Any]) -> None: global IN_ROOT, OUT_ROOT, EXCLUDE_VAL_FOLDER, START_DATE, END_DATE, DATE_PAIRS global DT_K, TU_K, TU_ITT, CONNECTIVITY, MIN_PIXELS_KEEP global USE_SIMPLE_CLOUD_MASK, USE_BTD_MASK, BTD2_GROW global SEED_WIN, SEED_MERGE_PIX, SEED_MIN_PIX, PROFILE_TIME IN_ROOT = _resolve_data_root(_resolve_path(cfg["input_root"])) OUT_ROOT = _resolve_path(cfg["output_root"]) os.makedirs(OUT_ROOT, exist_ok=True) EXCLUDE_VAL_FOLDER = bool(cfg.get("exclude_val_folder", EXCLUDE_VAL_FOLDER)) START_DATE = cfg.get("start_date", START_DATE) END_DATE = cfg.get("end_date", END_DATE) DATE_PAIRS = [tuple(pair) for pair in cfg.get("date_pairs", [])] DT_K = float(cfg.get("dt_k", DT_K)) TU_K = float(cfg.get("tu_k", TU_K)) TU_ITT = float(cfg.get("tu_itt", TU_ITT)) CONNECTIVITY = int(cfg.get("connectivity", CONNECTIVITY)) MIN_PIXELS_KEEP = int(cfg.get("min_pixels_keep", MIN_PIXELS_KEEP)) USE_SIMPLE_CLOUD_MASK = bool(cfg.get("use_simple_cloud_mask", USE_SIMPLE_CLOUD_MASK)) USE_BTD_MASK = bool(cfg.get("use_btd_mask", USE_BTD_MASK)) BTD2_GROW = float(cfg.get("btd2_grow", BTD2_GROW)) SEED_WIN = int(cfg.get("seed_win", SEED_WIN)) SEED_MERGE_PIX = int(cfg.get("seed_merge_pix", SEED_MERGE_PIX)) SEED_MIN_PIX = int(cfg.get("seed_min_pix", SEED_MIN_PIX)) PROFILE_TIME = bool(cfg.get("profile_time", PROFILE_TIME)) default_log = os.path.join(OUT_ROOT, "processing_errors.log") _set_logging(_resolve_path(cfg.get("log_file", default_log))) # ───────────── 저장 ───────────── def save_outputs(seg_map: np.ndarray, clusters: dict, out_prefix: str) -> None: """NetCDF + Pickle + static R-tree""" nc_path = out_prefix + ".nc" # 혹시 모를 찌꺼기 파일 삭제 if os.path.exists(nc_path): try: os.remove(nc_path) except OSError: pass # 1. 명시적으로 Dataset 객체 생성 ds = xr.DataArray( seg_map, dims=("r", "c"), coords={"r": np.arange(seg_map.shape[0]), "c": np.arange(seg_map.shape[1])}, name="region_id" ).to_dataset() # 2. engine을 'h5netcdf'로 변경 (세그멘테이션 폴트 회피의 핵심) ds.to_netcdf( nc_path, format="NETCDF4", engine="netcdf4", encoding={ "region_id": { "dtype": "float32", "zlib": True, "complevel": 9, "_FillValue": np.nan } } ) # 3. 파일 핸들을 명시적으로 닫아서 메모리 누수 및 open object 에러 방지 ds.close() # ---- Pickle 및 R-tree 저장 로직 ---- with open(out_prefix + "_clusters.pkl", "wb") as f: pickle.dump(clusters, f) def _cid_to_int(cid_str: str) -> int: try: if "_" in cid_str: return int(cid_str.split("_")[-1]) return int(cid_str) except Exception: return abs(hash(cid_str)) % (1 << 31) bulk = ((_cid_to_int(cid), tuple(v["bbox"]), None) for cid, v in clusters.items()) idx = rtree_index.Index(out_prefix + "_rtree", bulk) idx.close() # ───────────── 날짜/폴더 선택 ───────────── def is_date_in_range(timestamp: str, start_date: Optional[str], end_date: Optional[str]) -> bool: def compare_value(bound: Optional[str]) -> str: if bound is None or len(bound) <= 8: return timestamp[:8] return timestamp[:len(bound)] if start_date is not None and compare_value(start_date) < start_date: return False if end_date is not None and compare_value(end_date) > end_date: return False return True def is_timestamp_in_date_pairs(timestamp: str, date_pairs: List[Tuple[str, str]]) -> bool: for s, e in date_pairs: if is_date_in_range(timestamp, s, e): return True return False def get_target_folders_from_pairs(root_dir: str, date_pairs: List[Tuple[str, str]]) -> List[str]: ordered: List[str] = [] all_items = os.listdir(root_dir) for start_date, end_date in date_pairs: pair = [] for item in all_items: item_path = os.path.join(root_dir, item) if os.path.isdir(item_path) and len(item) == 8 and item.isdigit(): if EXCLUDE_VAL_FOLDER and "val" in item: continue if is_date_in_range(item + "0000", start_date, end_date): pair.append(item_path) pair.sort() ordered.extend(pair) return ordered def get_target_folders(root_dir: str, start_date: Optional[str], end_date: Optional[str]) -> List[str]: out: List[str] = [] for item in os.listdir(root_dir): item_path = os.path.join(root_dir, item) if os.path.isdir(item_path) and len(item) == 8 and item.isdigit(): if EXCLUDE_VAL_FOLDER and "val" in item: continue if is_date_in_range(item + "0000", start_date, end_date): out.append(item_path) out.sort() return out # ───────────── cloud mask ───────────── def get_cloud_mask(tb105: np.ndarray, btd_063_105: np.ndarray) -> np.ndarray: m = np.isfinite(tb105) if USE_SIMPLE_CLOUD_MASK: m &= (tb105 <= TU_K) if USE_BTD_MASK: m &= (btd_063_105 > BTD2_GROW) return m # ───────────── ITT Step1: seeds (local minima) ───────────── def find_markers_from_cold_minima(tb: np.ndarray, cloud_mask: np.ndarray) -> np.ndarray: valid = cloud_mask & np.isfinite(tb) if not valid.any(): return np.zeros_like(tb, dtype=np.int32) big = float(np.nanmax(tb[valid]) + 50.0) tb_for_min = np.where(valid, tb, big) mn = ndi.minimum_filter(tb_for_min, size=SEED_WIN, mode="nearest") cores = valid & (tb_for_min <= mn + 1e-6) if not cores.any(): return np.zeros_like(tb, dtype=np.int32) lab0, n0 = ndi.label(cores, structure=ndi.generate_binary_structure(2, 2)) if n0 > 0: thin = np.zeros_like(cores, dtype=bool) for k in range(1, n0 + 1): m = (lab0 == k) if not m.any(): continue vals = tb_for_min[m] idx = np.argmin(vals) rr, cc = np.where(m) thin[rr[idx], cc[idx]] = True cores = thin if SEED_MERGE_PIX and SEED_MERGE_PIX > 0: struct8 = ndi.generate_binary_structure(2, 2) cores = ndi.binary_dilation(cores, structure=struct8, iterations=int(SEED_MERGE_PIX)) lab1, n1 = ndi.label(cores, structure=struct8) if n1 > 0: thin2 = np.zeros_like(cores, dtype=bool) for k in range(1, n1 + 1): m = (lab1 == k) if not m.any(): continue vals = tb_for_min[m] idx = np.argmin(vals) rr, cc = np.where(m) thin2[rr[idx], cc[idx]] = True cores = thin2 lab, n = ndi.label(cores, structure=ndi.generate_binary_structure(2, 1)) if n == 0: return lab.astype(np.int32) if SEED_MIN_PIX and SEED_MIN_PIX > 1: sizes = np.bincount(lab.ravel()) keep = sizes >= int(SEED_MIN_PIX) keep[0] = False lab = np.where(keep[lab], lab, 0).astype(np.int32) lab, _ = ndi.label(lab > 0, structure=ndi.generate_binary_structure(2, 1)) return lab.astype(np.int32) # ───────────── fallback segmentation: CCL ───────────── def segment_all_clouds(cloud_mask: np.ndarray) -> Tuple[np.ndarray, int]: conn = 1 if CONNECTIVITY == 4 else 2 structure = ndi.generate_binary_structure(2, conn) lab, n = ndi.label(cloud_mask.astype(bool), structure=structure) return lab.astype(np.int32), int(n) # ───────────── segmentation: TRUE ITT (optimized neighbor masks) ───────────── def segment_clouds_itt(tb: np.ndarray, cloud_mask: np.ndarray) -> Tuple[np.ndarray, int]: valid = cloud_mask & np.isfinite(tb) if not valid.any(): return np.zeros_like(tb, dtype=np.int32), 0 conn = 1 if CONNECTIVITY == 4 else 2 nbh_struct = ndi.generate_binary_structure(2, conn) seg = find_markers_from_cold_minima(tb, valid) K = int(seg.max()) if K == 0: return seg.astype(np.int32), 0 CT = np.full((K + 1,), np.inf, dtype=np.float32) for i in range(1, K + 1): m = (seg == i) & valid if m.any(): CT[i] = float(np.nanmin(tb[m])) Tmin = float(np.nanmin(tb[valid])) if TU_ITT <= Tmin: return seg.astype(np.int32), int(seg.max()) thresholds = np.arange(Tmin + DT_K, TU_ITT + 1e-6, DT_K, dtype=np.float32) H, W = tb.shape for THT in thresholds: adj = ndi.binary_dilation(seg > 0, structure=nbh_struct) cand_new = (seg == 0) & valid & (tb < THT) & (~adj) if cand_new.any(): new_lab, n_new = ndi.label(cand_new, structure=ndi.generate_binary_structure(2, 1)) if n_new > 0: CT = np.pad(CT, (0, n_new), constant_values=np.inf) for j in range(1, n_new + 1): K += 1 seg[new_lab == j] = K m = (seg == K) & valid CT[K] = float(np.nanmin(tb[m])) if m.any() else np.inf cand = (seg == 0) & valid & (tb < THT) if not cand.any(): continue adj = ndi.binary_dilation(seg > 0, structure=nbh_struct) start = cand & adj if not start.any(): continue q = deque(map(tuple, np.argwhere(start))) inq = np.zeros_like(seg, dtype=bool) inq[start] = True while q: r, c = q.popleft() inq[r, c] = False if seg[r, c] != 0: continue if not (valid[r, c] and tb[r, c] < THT): continue neigh_ids = [] if CONNECTIVITY == 4: if r > 0 and seg[r - 1, c] > 0: neigh_ids.append(int(seg[r - 1, c])) if r < H - 1 and seg[r + 1, c] > 0: neigh_ids.append(int(seg[r + 1, c])) if c > 0 and seg[r, c - 1] > 0: neigh_ids.append(int(seg[r, c - 1])) if c < W - 1 and seg[r, c + 1] > 0: neigh_ids.append(int(seg[r, c + 1])) else: r0 = max(0, r - 1); r1 = min(H - 1, r + 1) c0 = max(0, c - 1); c1 = min(W - 1, c + 1) blk = seg[r0:r1 + 1, c0:c1 + 1] vals = blk[blk > 0] if vals.size: neigh_ids = [int(x) for x in vals.tolist()] if not neigh_ids: continue if len(neigh_ids) == 1: chosen = neigh_ids[0] else: uniq = list(set(neigh_ids)) tpx = float(tb[r, c]) chosen = uniq[0] dmin = np.inf for uid in uniq: d = abs(tpx - float(CT[uid])) if d < dmin: dmin = d chosen = uid seg[r, c] = chosen r0 = max(0, r - 1); r1 = min(H - 1, r + 1) c0 = max(0, c - 1); c1 = min(W - 1, c + 1) for rr in range(r0, r1 + 1): for cc in range(c0, c1 + 1): if rr == r and cc == c: continue if seg[rr, cc] == 0 and (not inq[rr, cc]) and valid[rr, cc] and (tb[rr, cc] < THT): q.append((rr, cc)) inq[rr, cc] = True return seg.astype(np.int32), int(seg.max()) # ───────────── seg -> clusters + seg_map_out ───────────── def seg_to_clusters(seg: np.ndarray, tstamp: str) -> Tuple[Dict[str, Any], Optional[np.ndarray]]: clusters: Dict[str, Any] = {} H, W = seg.shape ids = np.unique(seg) ids = ids[ids > 0] new_id = 0 for old in ids: mask = (seg == old) pix = int(mask.sum()) if pix < MIN_PIXELS_KEEP: continue rr, cc = np.where(mask) r1, r2 = int(rr.min()), int(rr.max()) c1, c2 = int(cc.min()), int(cc.max()) new_id += 1 cid_str = f"{tstamp}_{new_id}" clusters[cid_str] = { "bbox": [c1, r1, c2, r2], "pixel_count": pix, "flat_idx": np.flatnonzero(mask).astype(np.int32), "shape": [H, W], } if new_id == 0: return {}, None remap = np.zeros(int(seg.max()) + 1, dtype=np.int32) ni = 0 for old in ids: mask = (seg == old) if int(mask.sum()) < MIN_PIXELS_KEEP: continue ni += 1 remap[int(old)] = ni seg_int = remap[seg] seg_map_out = np.where(seg_int > 0, seg_int.astype(np.float32), np.nan).astype(np.float32) return clusters, seg_map_out # ───────────── per-file ───────────── def process_file(npy_path: str, rel_dir: str) -> bool: fn = os.path.basename(npy_path) tstamp = fn.split("_")[-1].replace(".npy", "") # 202309220830 if DATE_PAIRS: if not is_timestamp_in_date_pairs(tstamp, DATE_PAIRS): return False else: if not is_date_in_range(tstamp, START_DATE, END_DATE): return False out_dir = os.path.join(OUT_ROOT, rel_dir) os.makedirs(out_dir, exist_ok=True) out_pre = os.path.join(out_dir, fn[:-4]) # .npy 제거 (4글자) # ========================================================= # 이미 결과 파일(.nc)이 존재하면 건너뛰기 if os.path.exists(out_pre + ".nc"): # 필요시 print 주석 해제하여 로그 확인 print(f"[건너뛰기] 이미 처리된 파일: {fn}") return True # 이미 성공한 것으로 간주하여 True 반환 # ========================================================= t_total0 = time.perf_counter() # ---- IO read (NPY 파일 처리) ---- t_io_read0 = time.perf_counter() try: # npy 데이터를 메모리에 로드 data = np.load(npy_path, allow_pickle=True) # [주의] 데이터 저장 방식에 따라 아래 코드를 적절히 선택해야 합니다. # 1) npy 파일이 딕셔너리로 저장된 경우 (예: np.save('...', {'ir105': arr1, 'wv063': arr2})) if data.dtype == object and isinstance(data.item(), dict): data_dict = data.item() if "ir105" not in data_dict or "wv063" not in data_dict: logger.error(f"{rel_dir}/{fn} missing keys in dict") return False bt105 = data_dict["ir105"].astype(np.float32) bt063 = data_dict["wv063"].astype(np.float32) # 2) npy 파일이 (2, H, W) 형태의 다차원 배열로 저장된 경우 # (예: 0번 채널이 ir105, 1번 채널이 wv063) else: # 형태가 다를 경우 인덱스 [0], [1]을 데이터 구조에 맞게 변경하세요. bt105 = data[0].astype(np.float32) bt063 = data[1].astype(np.float32) except Exception as e: logger.error(f"{rel_dir}/{fn} read error: {e}") return False t_io_read = time.perf_counter() - t_io_read0 # ---- CODE ---- t_code0 = time.perf_counter() btd_063_105 = bt063 - bt105 cloud_mask = get_cloud_mask(bt105, btd_063_105) seg, nseg = segment_clouds_itt(bt105, cloud_mask) if nseg < 1: seg, _ = segment_all_clouds(cloud_mask) clusters, seg_map_out = seg_to_clusters(seg, tstamp) if not clusters or seg_map_out is None: return False # ---- IO save ---- t_io_save0 = time.perf_counter() save_outputs(seg_map_out, clusters, out_pre) t_io_save = time.perf_counter() - t_io_save0 t_code = time.perf_counter() - t_code0 t_total = time.perf_counter() - t_total0 if PROFILE_TIME: print( f"✔ seg-only {rel_dir}/{fn} (N={len(clusters)}) | " f"IO(read)={t_io_read:.2f}s, ITT={t_code:.2f}s, IO(save)={t_io_save:.2f}s | " f"total={t_total:.2f}s" ) return True # ───────────── main ───────────── if __name__ == "__main__": parser = argparse.ArgumentParser(description="Build per-scene cloud objects by region growing.") parser.add_argument("--config", required=True, help="JSON configuration path") parser.add_argument("--device", default=None, help="Accepted for a common CLI; this stage runs on CPU") parser.add_argument("--output-dir", default=None, help="Override output_root") args = parser.parse_args() config_path = os.path.abspath(args.config) RUN_DIR = os.path.dirname(config_path) cfg = load_config(config_path) if args.output_dir: cfg["output_root"] = os.path.abspath(args.output_dir) apply_config(cfg) t0 = datetime.now() total_files = 0 ok_files = 0 total_time = 0.0 print("\n" + "=" * 70) print("SEGMENTATION ONLY (cloud_mask + TRUE ITT(Tu=285K) + save seg/rtree)") if DATE_PAIRS: target_folders = get_target_folders_from_pairs(IN_ROOT, DATE_PAIRS) else: target_folders = get_target_folders(IN_ROOT, START_DATE, END_DATE) print("=" * 70) if not target_folders: print("처리할 폴더가 없습니다.") raise SystemExit(0) for folder_path in target_folders: rel = os.path.relpath(folder_path, IN_ROOT) try: # .nc 대신 .npy 확장자를 찾도록 수정 npys = sorted(f for f in os.listdir(folder_path) if f.endswith(".npy")) except OSError as e: print(f"[에러] 폴더 읽기 실패: {rel} - {e}") continue if not npys: print(f"[건너뛰기] npy 파일 없음: {rel}") continue print(f"[DAY] 폴더: {rel} ({len(npys)}개)") for f in npys: file_start = time.perf_counter() ok = process_file(os.path.join(folder_path, f), rel) file_time = time.perf_counter() - file_start total_files += 1 total_time += file_time if ok: ok_files += 1 avg_time = total_time / max(total_files, 1) print(f"└─ 파일 처리 시간: {file_time:.1f}초 (평균: {avg_time:.1f}초)") total_elapsed = (datetime.now() - t0).total_seconds() print("\n" + "=" * 70) print("처리 완료!") print(f"총 시도 파일 수: {total_files}개") print(f"성공 파일 수: {ok_files}개") print(f"총 소요 시간: {total_elapsed:.1f}초") print(f"파일당 평균 처리 시간: {total_time / max(total_files, 1):.1f}초") print("=" * 70 + "\n")