#!/usr/bin/env python3 # -*- coding: utf-8 -*- """ Temporal Overlap Tracking (static R-tree, backward link 저장) 개선 개요 - step1의 region_id.nc + *_clusters.pkl + *_rtree 인덱스를 읽어 시간 역추적을 수행한다. - 현재 시각의 강수 후보 클러스터를 seed로 두고, 과거 클러스터와 픽셀 overlap을 계산해 동일 생애주기를 연결한다. - 큰 클러스터/작은 클러스터 비율, 최소 overlap 비율, 최대 backtracking 시간으로 과연결을 제한한다. - 최신 버전은 위경도/낙뢰 보조파일을 사용하지 않고, BT 입력의 hsr 필드와 step1 결과만 사용한다. * 입력 - step1 region 결과: region_id.nc, *_clusters.pkl, *_rtree.{idx,dat} - BT/RADAR npy: concat_gk2a_radar_YYYYMMDDHHMM.npy (hsr 사용) * 결과물 - *_label.nc: temporal overlap으로 연결된 label - *_visited.pkl: 처리한 클러스터 방문 상태 - *_links.pkl: {"현재시각_id": ["과거시각_id", ...]} 링크 정보 경로 설정 - REGION_ROOT/BT_ROOT/OUT_ROOT는 import 시점에는 빈 문자열이다. - main에서 build/run/step2_temporal_overlapping_config.json을 읽은 뒤 apply_config()가 실제 경로로 채운다. - bt_root가 data_preprocess/result/res_2km처럼 L1B/L2 폴더이면, 실제 labeling 입력인 L1B/YYYYMMDD를 자동 선택한다. """ from __future__ import annotations import argparse import os, sys, pickle, time, gc from datetime import datetime, timedelta from collections import deque from typing import Dict, Tuple, List, Optional import numpy as np import xarray as xr from rtree import index from tqdm import tqdm # ─────────── 경로 / 파라미터 ────────────────────────────────────── try: sys.stdout.reconfigure(encoding="utf-8", errors="replace") sys.stderr.reconfigure(encoding="utf-8", errors="replace") except AttributeError: pass CONFIG_NAME = "step2_temporal_overlapping_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 # Config 적용 전 placeholder. 실제 값은 apply_config()에서 채운다. REGION_ROOT = "" BT_ROOT = "" OUT_ROOT = "" DBZ_THR = 35.0 STEP_MIN = 10 MAX_BACK_MINUTES = 120 MIN_OVERLAP_RATIO = 0.4 SIZE_RATIO_SMALL = 0.1 SIZE_RATIO_LARGE = 1.2 MAX_BACKTRACK_PIXELS = 10000 START_DATE = '202508010000' END_DATE = '202510312350' def apply_config(cfg: dict) -> None: global REGION_ROOT, BT_ROOT, OUT_ROOT, DBZ_THR, STEP_MIN, MAX_BACK_MINUTES global MIN_OVERLAP_RATIO, SIZE_RATIO_SMALL, SIZE_RATIO_LARGE, MAX_BACKTRACK_PIXELS global START_DATE, END_DATE REGION_ROOT = _resolve_path(cfg["region_root"]) BT_ROOT = _resolve_data_root(_resolve_path(cfg["bt_root"])) OUT_ROOT = _resolve_path(cfg["output_dir"]) os.makedirs(OUT_ROOT, exist_ok=True) DBZ_THR = float(cfg.get("dbz_thr", DBZ_THR)) STEP_MIN = int(cfg.get("step_min", STEP_MIN)) MAX_BACK_MINUTES = int(cfg.get("max_back_minutes", MAX_BACK_MINUTES)) MIN_OVERLAP_RATIO = float(cfg.get("min_overlap_ratio", MIN_OVERLAP_RATIO)) SIZE_RATIO_SMALL = float(cfg.get("size_ratio_small", SIZE_RATIO_SMALL)) SIZE_RATIO_LARGE = float(cfg.get("size_ratio_large", SIZE_RATIO_LARGE)) MAX_BACKTRACK_PIXELS = int(cfg.get("max_backtrack_pixels", MAX_BACKTRACK_PIXELS)) date_pairs = cfg.get("date_pairs") or [] if date_pairs: START_DATE = date_pairs[0][0] END_DATE = date_pairs[-1][1] else: START_DATE = cfg.get("start_date", START_DATE) END_DATE = cfg.get("end_date", END_DATE) # ─────────── 패키지( seg / clusters / rtree ) LRU 캐시 ─────────── # seg는 int32 라벨맵으로 변환해서 캐시 (0=배경) _pkg_cache : dict[str, Tuple[np.ndarray, Dict[str,dict], index.Index]] = {} _hsr_cache: dict[str, np.ndarray] = {} # ─────────── ROI mask 유틸 (핵심) ──────────────────────────────── def _bbox_slices(bbox: List[int]) -> Tuple[slice, slice]: c1, r1, c2, r2 = bbox return slice(r1, r2 + 1), slice(c1, c2 + 1) def roi_mask_from_seg(seg_i: np.ndarray, cid_int: int, bbox: List[int]) -> Tuple[slice, slice, np.ndarray]: """ seg_i(int32, 0=배경) + bbox로 ROI에서만 mask 생성 (벡터화, 파이썬 for 없음) 반환: (rsl, csl, mask_roi[bool]) """ rsl, csl = _bbox_slices(bbox) mask_roi = (seg_i[rsl, csl] == cid_int) return rsl, csl, mask_roi def parse_cid_int(cid_str: str) -> int: # "YYYYMMDDHHMM_N" -> N return int(cid_str.split('_')[-1]) # ─────────── I/O: seg/clusters/rtree 로드 ──────────────────────── def load_pkg(ts: str) -> Tuple[np.ndarray, Dict[str, dict], index.Index] | None: """seg_map(int32), clusters(dict), rtree (or None)""" if ts in _pkg_cache: return _pkg_cache[ts] day = ts[:8] pref = os.path.join(REGION_ROOT, day, f"concat_gk2a_radar_{ts}") nc = pref + ".nc" pkl = pref + "_clusters.pkl" idxf = pref + "_rtree" if not (os.path.exists(nc) and os.path.exists(pkl) and os.path.exists(idxf + ".idx")): return None # region_id: float32, NaN background -> int32, 0 background seg_f = xr.open_dataset(nc)["region_id"].values # float32 (NaN background) seg_i = np.where(np.isnan(seg_f), 0, seg_f).astype(np.int32) with open(pkl, "rb") as f: clusters = pickle.load(f) rtree_idx = index.Index(idxf) _pkg_cache[ts] = (seg_i, clusters, rtree_idx) return _pkg_cache[ts] def load_hsr(ts: str) -> Optional[np.ndarray]: """concat_gk2a_radar_{ts}.npy에서 HSR 필드 로드 (dict['hsr']) 캐시""" if ts in _hsr_cache: return _hsr_cache[ts] npy_bt = os.path.join(BT_ROOT, ts[:8], f"concat_gk2a_radar_{ts}.npy") if not os.path.exists(npy_bt): raise FileNotFoundError(f"HSR file missing: {npy_bt}") # --- 파일 손상 예외 처리 --- try: data = np.load(npy_bt, allow_pickle=True).item() except (EOFError, pickle.UnpicklingError) as e: print(f"[경고] 파일 손상으로 로드 실패 (건너뜀): {npy_bt} | 에러: {e}") return None except Exception as e: print(f"[경고] 예상치 못한 에러로 파일 로드 실패: {npy_bt} | 에러: {e}") return None if not isinstance(data, dict): raise ValueError(f"BT npy 로드 결과가 dict가 아닙니다: {type(data)} ({npy_bt})") if "hsr" not in data: raise KeyError(f"'hsr' key missing in BT npy: {npy_bt}") hsr = data["hsr"] _hsr_cache[ts] = hsr return hsr # ─────────── truth / visited / links 캐시+I/O ──────────────────── _truth_cache : dict[str, np.ndarray] = {} _vis_cache : dict[str, dict[str, bool]] = {} _link_cache : dict[str, Dict[str, List[str]]] = {} def _path(kind: str, ts: str) -> str: day = ts[:8] if kind == "truth": return os.path.join(OUT_ROOT, day, f"{ts}_label.nc") elif kind == "visited": return os.path.join(OUT_ROOT, day, f"{ts}_visited.pkl") elif kind == "links": return os.path.join(OUT_ROOT, day, f"{ts}_links.pkl") else: raise ValueError def load_truth(ts: str, shape: Tuple[int, int]) -> np.ndarray: if ts in _truth_cache: return _truth_cache[ts] p = _path("truth", ts) if os.path.exists(p): _truth_cache[ts] = xr.open_dataset(p)["label"].values else: _truth_cache[ts] = np.full(shape, np.nan, dtype=np.float32) return _truth_cache[ts] def save_truth(ts: str): if ts not in _truth_cache: return arr = _truth_cache[ts] output_path = _path("truth", ts) day = ts[:8] os.makedirs(os.path.join(OUT_ROOT, day), exist_ok=True) da = xr.DataArray( arr, dims=("r", "c"), coords={"r": np.arange(arr.shape[0]), "c": np.arange(arr.shape[1])}, name="label", ) da.to_dataset().to_netcdf( output_path, format="NETCDF4", encoding={"label": {"dtype": "float32", "zlib": True, "complevel": 9}}, ) def load_visited(ts: str) -> dict[str, bool]: if ts in _vis_cache: return _vis_cache[ts] p = _path("visited", ts) if os.path.exists(p): with open(p, "rb") as f: _vis_cache[ts] = pickle.load(f) else: _vis_cache[ts] = {} return _vis_cache[ts] def save_visited(ts: str): if ts not in _vis_cache: return output_path = _path("visited", ts) day = ts[:8] os.makedirs(os.path.join(OUT_ROOT, day), exist_ok=True) with open(output_path, "wb") as f: pickle.dump(_vis_cache[ts], f, pickle.HIGHEST_PROTOCOL) def load_links(ts: str) -> Dict[str, List[str]]: if ts in _link_cache: return _link_cache[ts] p = _path("links", ts) if os.path.exists(p): with open(p, "rb") as f: _link_cache[ts] = pickle.load(f) else: _link_cache[ts] = {} return _link_cache[ts] def save_links(ts: str): if ts not in _link_cache: return output_path = _path("links", ts) day = ts[:8] os.makedirs(os.path.join(OUT_ROOT, day), exist_ok=True) with open(output_path, "wb") as f: pickle.dump(_link_cache[ts], f, pickle.HIGHEST_PROTOCOL) # ─────────── 캐시 메모리 관리 ───────────────────────────────────── def cleanup_cache(current_ts: str): try: current_dt = datetime.strptime(current_ts, "%Y%m%d%H%M") cutoff_dt = current_dt - timedelta(minutes=MAX_BACK_MINUTES + 60) cutoff_ts = cutoff_dt.strftime("%Y%m%d%H%M") total_removed = 0 to_remove = [ts for ts in _pkg_cache.keys() if ts < cutoff_ts] for ts in to_remove: _, _, rtree_idx = _pkg_cache[ts] try: rtree_idx.close() except Exception: pass del _pkg_cache[ts] total_removed += len(to_remove) for cache in (_truth_cache, _vis_cache, _link_cache, _hsr_cache): to_remove = [ts for ts in cache.keys() if ts < cutoff_ts] for ts in to_remove: del cache[ts] total_removed += len(to_remove) if total_removed > 0: gc.collect() print(f"[캐시 정리] {total_removed}개 시간 데이터 제거 (cutoff: {cutoff_ts})") except Exception as e: print(f"[캐시 정리 오류] {e}") # ─────────── 성숙도 판별 (mask 없이: seg+bbox ROI) ───────────────── def check_cluster_maturity(ts: str, seg_i: np.ndarray, cid_int: int, bbox: List[int]) -> int | None: """ HSR에서 클러스터 내부 최대값으로 성숙도 판별. mask는 seg_i+bbox ROI에서 생성. Returns: 1(성숙: >=35dBZ), 2(미성숙), None(HSR 파일/데이터 없음) """ try: hsr_full = load_hsr(ts) except (FileNotFoundError, KeyError, ValueError): return None if hsr_full is None: return None rsl, csl, mask_roi = roi_mask_from_seg(seg_i, cid_int, bbox) if not mask_roi.any(): return 2 hsr_roi = hsr_full[rsl, csl] vals = hsr_roi[mask_roi] vals = vals[~np.isnan(vals)] if vals.size == 0: return 2 return 1 if float(vals.max()) >= DBZ_THR else 2 # ─────────── 메인 처리 (한 타임스텝) ─────────────────────────────── def process_ts(ts: str) -> float: t0 = time.time() cleanup_cache(ts) pkg_now = load_pkg(ts) if pkg_now is None: print(f"[skip] {ts} (no package)") return 0.0 seg_now, clusters_now, _ = pkg_now h, w = seg_now.shape # ---- 1단계: HSR 임계값으로 rainy cluster 후보 cid 수집 --------- try: hsr = load_hsr(ts) except (FileNotFoundError, KeyError, ValueError): print(f"[skip] {ts} (no HSR or 'hsr' key for seed generation)") return 0.0 # hsr이 None인 경우 (파일 손상 등) 메인루프 에러를 막기 위해 0.0 반환 if hsr is None: print(f"[알림] {ts} 시각의 HSR 데이터 손상으로 처리를 건너뜁니다.") return 0.0 rs, cs = np.where(hsr >= DBZ_THR) cappi_seed_cids = set() for r, c in zip(rs, cs): cid = int(seg_now[r, c]) # seg_now는 int32, 0=배경 if cid > 0: cappi_seed_cids.add(cid) if not cappi_seed_cids: print(f"[skip] {ts} (no rainy cluster)") return 0.0 # ---- 2단계: 강수 클러스터를 바로 seed로 사용 (낙뢰 필터링 제거) ---- seed_cids = cappi_seed_cids # ---- cache 객체 로드 ------------------------------------------ load_truth(ts, (h, w)) load_visited(ts) links_now = load_links(ts) Q = deque() # cid_str only def enqueue_cluster(cid_str: str, cid_int: int, seg_ts: np.ndarray, bbox: List[int], is_seed: bool = False): """truth/visited/큐 관리 (mask 없이 ROI에서만 truth 채움)""" ts_local = cid_str.split('_')[0] vdict = load_visited(ts_local) if cid_str in vdict: return vdict[cid_str] = is_seed truth = load_truth(ts_local, (h, w)) rsl, csl, mask_roi = roi_mask_from_seg(seg_ts, cid_int, bbox) # ROI view에 직접 할당 truth_roi = truth[rsl, csl] truth_roi[mask_roi] = cid_int Q.append(cid_str) # ---- seed enqueue --------------------------------------------- for cid_int in seed_cids: cid_str = f"{ts}_{cid_int}" if cid_str not in clusters_now: continue info = clusters_now[cid_str] # MAX_BACKTRACK_PIXELS를 넘는 클러스터는 큐에 넣지 않음 if info.get("pixel_count", 0) > MAX_BACKTRACK_PIXELS: continue enqueue_cluster(cid_str, cid_int, seg_now, info["bbox"], is_seed=True) links_now.setdefault(cid_str, []) ts_dt = datetime.strptime(ts, "%Y%m%d%H%M") # ---- 시간 역방향 BFS ------------------------------------------ while Q: cur_cid = Q.popleft() cur_ts = cur_cid.split('_')[0] cur_int = parse_cid_int(cur_cid) pkg_cur = load_pkg(cur_ts) if pkg_cur is None: continue seg_cur, clusters_cur, _ = pkg_cur if cur_cid not in clusters_cur: continue cur_info = clusters_cur[cur_cid] c1, r1, c2, r2 = cur_info["bbox"] # ts - cur_ts 차이가 너무 크면 stop cur_dt = datetime.strptime(cur_ts, "%Y%m%d%H%M") if (ts_dt - cur_dt).total_seconds() / 60 > MAX_BACK_MINUTES: link_dict = load_links(cur_ts) link_dict.setdefault(cur_cid, []) save_truth(cur_ts); save_visited(cur_ts); save_links(cur_ts) continue prev_dt = cur_dt - timedelta(minutes=STEP_MIN) prev_ts = prev_dt.strftime("%Y%m%d%H%M") visited_prev = load_visited(prev_ts) overlaps: List[str] = [] def process_prev_clusters(pkg_prev) -> List[str]: """이전 시점 클러스터들 처리 (mask 없이 seg ROI 비교로 overlap 계산)""" if pkg_prev is None: return [] seg_prev, clusters_prev, rtree_prev = pkg_prev local_overlaps: List[str] = [] # bbox 교차 후보만 R-tree로 가져옴 for prev_int in rtree_prev.intersection((c1, r1, c2, r2)): prev_cid = f"{prev_ts}_{prev_int}" if prev_cid not in clusters_prev: continue prev_info = clusters_prev[prev_cid] p1, q1, p2, q2 = prev_info["bbox"] ic1 = max(c1, p1); ir1 = max(r1, q1) ic2 = min(c2, p2); ir2 = min(r2, q2) if ic1 > ic2 or ir1 > ir2: continue # overlap: seg 값 비교로 바로 계산 (mask 만들 필요 없음) cur_roi = seg_cur[ir1:ir2+1, ic1:ic2+1] prev_roi = seg_prev[ir1:ir2+1, ic1:ic2+1] overlap_count = np.count_nonzero((cur_roi == cur_int) & (prev_roi == prev_int)) if overlap_count == 0: continue # 조건 1: Overlap ratio (overlap / prev_cluster_pixels) prev_pix = prev_info.get("pixel_count", 0) if prev_pix <= 0: continue if (overlap_count / prev_pix) < MIN_OVERLAP_RATIO: continue # 조건 2/3: 크기 비율 필터 cur_pix = cur_info.get("pixel_count", 0) if cur_pix <= 0: continue if prev_pix < SIZE_RATIO_SMALL * cur_pix: continue if prev_pix > SIZE_RATIO_LARGE * cur_pix: continue local_overlaps.append(prev_cid) # 아직 방문 안 되었고 truth에서 비어있으면 큐에 추가 if prev_cid not in visited_prev: prev_truth_map = load_truth(prev_ts, (h, w)) # prev 클러스터 ROI mask로 기존 truth 채워짐 여부 체크 rsl, csl, mask_roi = roi_mask_from_seg(seg_prev, prev_int, prev_info["bbox"]) if mask_roi.any(): existing_values = prev_truth_map[rsl, csl][mask_roi] if not np.any(~np.isnan(existing_values)): enqueue_cluster(prev_cid, prev_int, seg_prev, prev_info["bbox"], is_seed=False) return local_overlaps pkg_prev = load_pkg(prev_ts) overlaps = process_prev_clusters(pkg_prev) # ---- 링크 누적 저장 -------------------------------------- link_dict = load_links(cur_ts) link_dict.setdefault(cur_cid, []) link_dict[cur_cid] = list(set(link_dict[cur_cid]) | set(overlaps)) save_truth(cur_ts); save_visited(cur_ts); save_links(cur_ts) print(f"✔ {ts} (total visited {len(load_visited(ts))})") return time.time() - t0 # ─────────── 전체 날짜 순회 ─────────────────────────────────────── if __name__ == "__main__": parser = argparse.ArgumentParser(description="Link cloud objects through time.") 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_dir"] = os.path.abspath(args.output_dir) apply_config(cfg) nc_files: List[str] = [] for root, _, fs in os.walk(REGION_ROOT): nc_files += [ os.path.join(root, f) for f in fs if f.endswith(".nc") and "concat_gk2a_radar_" in f ] nc_files.sort(key=lambda p: os.path.basename(p)[-15:-3]) # 날짜 필터링 if START_DATE is not None or END_DATE is not None: original_count = len(nc_files) filtered_files = [] for f in nc_files: ts = os.path.basename(f)[-15:-3] if START_DATE is not None and ts < START_DATE: continue if END_DATE is not None and ts > END_DATE: break filtered_files.append(f) nc_files = filtered_files print(f"\n날짜 필터링: {original_count}개 → {len(nc_files)}개 파일") if START_DATE: print(f"시작 날짜: {START_DATE}") if END_DATE: print(f"종료 날짜: {END_DATE}") if not nc_files: print("처리할 파일이 없습니다.") raise SystemExit(0) total_time = 0.0 total_files = 0 t0 = datetime.now() print(f"\n총 {len(nc_files)}개 파일 처리 시작") print("=" * 50) # ─────────── 메인 루프 ──────────────────────────────────────── for ncp in tqdm(nc_files, desc="TIMESTEPS"): ts = os.path.basename(ncp)[-15:-3] file_time = process_ts(ts) # NoneType 에러 방지용 방어 코드 추가 if file_time is None: continue if file_time > 0: total_time += file_time total_files += 1 avg_time = total_time / total_files print(f"└─ 파일 처리 시간: {file_time:.1f}초 (평균: {avg_time:.1f}초)") total_elapsed = (datetime.now() - t0).total_seconds() print("\n" + "=" * 50) print("처리 완료!") print(f"총 처리 파일 수: {total_files}개") print(f"총 소요 시간: {total_elapsed:.1f}초") if total_files > 0: print(f"파일당 평균 처리 시간: {total_time / total_files:.1f}초") print("=" * 50 + "\n")