From 53a3acd52d1301e971b6f90982fea73100dc2562 Mon Sep 17 00:00:00 2001 From: Yijing-Zuo Date: Fri, 26 Sep 2025 12:39:02 -0500 Subject: [PATCH 01/19] feat(ms): add multi-snapshot sampler + selfcheck test --- inc/samp_ms.py | 192 +++++++++++++++++++++++++++++++++++++ tests/test_ms_selfcheck.py | 50 ++++++++++ 2 files changed, 242 insertions(+) create mode 100644 inc/samp_ms.py create mode 100644 tests/test_ms_selfcheck.py diff --git a/inc/samp_ms.py b/inc/samp_ms.py new file mode 100644 index 0000000..af6fbe5 --- /dev/null +++ b/inc/samp_ms.py @@ -0,0 +1,192 @@ +# -*- coding: utf-8 -*- +""" +每一步先处理 R->I,再处理 I->S; +I->S: +按 zI 降序遍历,当某个时刻 t 有观测时,对相应节点的本步“反向转移”进行计算, +以确保 Y[t] 与观测一致;若与 SIR 反向可达性或容量约束冲突,则将该样本标记为 invalid; +若 compute_lik=True,则返回在 Qθ 下的 log-likelihood: + +R->I:对 (y_{t+1}==R) 的位置累加 lR(不变); +I->S:只对 **最终 msk_opt=True** 的位置累加 lI(与原版一致); +对于invalid 样本的 lik 置为 -inf + +原版:Y = samp_ms(q_net, y_T, zI, zR, n_samples=64, compute_lik=False) +现有版本:Y, lq = samp_ms( + q_net, y_T, zI, zR, n_samples=64, compute_lik=True, + obs_times=[3, 7], + obs_states=torch.stack([y3_obs, y7_obs]), # (2, n_nodes) Long + obs_masks=torch.stack([m3, m7]) # (2, n_nodes) Bool +) +""" + +import torch +try: + from .diffus import SIR_STATES +except Exception: + from types import SimpleNamespace + SIR_STATES = SimpleNamespace(S=0, I=1, R=2) +def torch_log(x: torch.Tensor) -> torch.Tensor: + eps = 1e-12 + return torch.log(torch.clamp(x, min=eps)) + +@torch.no_grad() +def samp_ms( + q_net, + y, + zI, + zR, + n_samples, + compute_lik=False, + obs_times=None,#list[int] / 1D LongTensor,观测时刻(取值{0..T};允许 T 表示 y_T) + obs_states=None,#(K, nodes) Long,obs_times 对应的观测状态 + obs_masks=None#(K, nodes) Bool,部分观测掩码;None 表示全 True +): + """ + 返回:若compute_lik=False -> Y: (T, nodes, samples) Long + 若compute_lik=True -> (Y, lik),其中 lik: (samples,) Float + """ + + device = q_net.device + T = int(q_net.T) + n_nodes = int(q_net.n_nodes) + obs_map = {} + if obs_times is not None: + if not torch.is_tensor(obs_times): + obs_times = torch.as_tensor(obs_times, dtype=torch.long, device=device) + else: + obs_times = obs_times.to(device=device, dtype=torch.long) + + assert obs_states is not None#obs_states 必须与 obs_times 同时提供 + if not torch.is_tensor(obs_states): + obs_states = torch.as_tensor(obs_states, dtype=torch.long, device=device) + else: + obs_states = obs_states.to(device=device, dtype=torch.long) + assert obs_states.dim() == 2 and obs_states.size(1) == n_nodes#obs_states 必须为 (K, nodes) + assert obs_states.size(0) == obs_times.numel()#obs_states 第一维必须等于 obs_times 的长度 + if obs_masks is None: + obs_masks = torch.ones_like(obs_states, dtype=torch.bool, device=device) + else: + if not torch.is_tensor(obs_masks): + obs_masks = torch.as_tensor(obs_masks, dtype=torch.bool, device=device) + else: + obs_masks = obs_masks.to(device=device, dtype=torch.bool) + assert obs_masks.shape == obs_states.shape#obs_masks 形状必须与 obs_states 一致 + + for k, t in enumerate(obs_times.tolist()): + y_obs_t = obs_states[k] + m_obs_t = obs_masks[k] + if t == T: + obs_map[T] = (y_obs_t, m_obs_t) + else: + assert 0 <= t < T#obs_times 的取值必须位于 [0, T] + obs_map[t] = (y_obs_t, m_obs_t) + zI, uid = zI.sort(dim=1, descending=True) + uid = uid.squeeze(dim=2) + qI = torch.sigmoid(zI) + xI = SIR_STATES.I - qI.expand(-1, -1, n_samples).bernoulli().long() + lI = torch_log(torch.where(xI != SIR_STATES.I, qI, 1. - qI)) + qR = torch.sigmoid(zR) + xR = SIR_STATES.R - qR.expand(-1, -1, n_samples).bernoulli().long() + lR = torch_log(torch.where(xR != SIR_STATES.R, qR, 1. - qR)) + y = y.to(device=device, dtype=torch.long).unsqueeze(1).expand(-1, n_samples)# (nodes, samples) + Y = torch.empty(T, n_nodes, n_samples, dtype=torch.long, device=device)# (T, nodes, samples) + + if compute_lik: + lik = q_net.zero + invalid = torch.zeros(n_samples, dtype=torch.bool, device=device) + if T in obs_map: + y_obs_T, m_obs_T = obs_map[T] + bad_T = m_obs_T.unsqueeze(1) & (y != y_obs_T.unsqueeze(1)) + if bad_T.any(): + invalid |= bad_T.any(dim=0) + for t in range(T - 1, -1, -1): + y_obs_t = None + m_obs_t = None + if t in obs_map: + y_obs_t, m_obs_t = obs_map[t] + # 反向可达性: + #y_{t+1}=S -> y_t 必为 S,否则 invalid + #y_{t+1}=I -> y_t 不可为 R,否则 invalid + #y_{t+1}=R -> y_t ∈ {R,I,S} 均可达(先 R->I,再可能 I->S) + bad_from_S = (y == SIR_STATES.S) & m_obs_t.unsqueeze(1) & (y_obs_t != SIR_STATES.S).unsqueeze(1) + bad_from_I = (y == SIR_STATES.I) & m_obs_t.unsqueeze(1) & (y_obs_t == SIR_STATES.R).unsqueeze(1) + bad_any = bad_from_S | bad_from_I + if bad_any.any(): + invalid |= bad_any.any(dim=0) + mR = (y == SIR_STATES.R) + if y_obs_t is not None: + mR_obs = mR & m_obs_t.unsqueeze(1) + qR_t = qR[t].expand(-1, n_samples) + # 观测想要 y_t=R -> 强制“不发生 R->I”,保持 R + wantR = mR_obs & (y_obs_t.eq(SIR_STATES.R).unsqueeze(1)) + if wantR.any(): + xR[t][wantR] = SIR_STATES.R + lR[t][wantR] = torch_log(1. - qR_t[wantR]) + # 观测想要 y_t ∈ {I,S} -> 强制“发生 R->I”,先变 I + wantI_or_S = mR_obs & (~y_obs_t.eq(SIR_STATES.R).unsqueeze(1)) + if wantI_or_S.any(): + xR[t][wantI_or_S] = SIR_STATES.I + lR[t][wantI_or_S] = torch_log(qR_t[wantI_or_S]) + # 若目标是 S,则稍后 I->S 阶段还会再强制一次(最终变 S) + y = torch.where(mR, xR[t], y) + if compute_lik: + lik = lik + torch.where(mR, lR[t], q_net.zero).sum(dim=0) + msk = (y == SIR_STATES.I)# (nodes, samples) + rem = torch.where(msk, q_net.rem, q_net.n_inf)# (nodes, samples) 广播 rem 初值 + for i, u in enumerate(uid[t]):# u: 节点 id(0..n_nodes-1) + if msk[u].max():# 该节点在任一样本为 I 才需要处理 + vid = q_net.neighbs[u.item()]# (deg_u,) Long,邻居索引列表 + # 原版 I->S 可行性:opt = (rem[u] > 1) & (min_{v∈N(u)} rem[v] > 1) + opt = (rem[u] > 1) & (rem[vid].min(dim=0).values > 1) # (samples,) + msk_opt = msk[u] & opt # (samples,) + # 若 t 时刻有观测,对该节点进行强制(仅当此节点被观测) + if (y_obs_t is not None) and bool(m_obs_t[u.item()]): + target = int(y_obs_t[u.item()].item()) + if target == SIR_STATES.S: + # 目标为 S -> 强制“发生 I->S” + need_I = ~msk[u] + if need_I.any(): + invalid |= need_I + # 覆盖本节点排序坐标 (t, i) 的采样与对数概率 + xI[t, i] = SIR_STATES.S + # qI[t, i] 形状为 (1,),扩展到 (samples,) + lI[t, i] = torch_log(qI[t, i].expand(n_samples)) + # 强制发生时仍需 opt 通过;否则该样本 invalid + bad_no_opt = msk[u] & (~opt) + if bad_no_opt.any(): + invalid |= bad_no_opt + elif target == SIR_STATES.I: + # 目标为 I -> 强制“不发生 I->S” + need_I = ~msk[u] + if need_I.any(): + invalid |= need_I + xI[t, i] = SIR_STATES.I + lI[t, i] = torch_log(1. - qI[t, i].expand(n_samples)) + # 不发生时,无需 opt(与原版一致:opt=False 时不会进入 toss) + else: + # 目标为 R:本步(I->S 阶段)无法直接达到 R,不可达 + invalid |= msk[u] + #应用 I->S 到 y(仅在 msk_opt 的样本上生效) + y[u] = torch.where(msk_opt, xI[t, i], y[u]) + trs = (y[u] != SIR_STATES.I)# 是否发生了 I->S(= 正向新感染) + + rem[u] = torch.where(msk[u], torch.where(trs, rem[u] - 1, q_net.n_inf), rem[u]) + rem[vid] = torch.where(msk[u].unsqueeze(0), + torch.where(trs.unsqueeze(0), rem[vid] - 1, q_net.n_inf), + rem[vid]) + #最终 msk[u] 仅保留“进入 toss 且 opt 通过”的样本(与原版一致) + msk[u] = msk_opt + Y[t] = y + if compute_lik: + lik = lik + torch.where(msk[uid[t]], lI[t], q_net.zero).sum(dim=0) # (samples,) + + Y = Y.detach().clone() + if compute_lik: + neg_inf = torch.full_like(lik if torch.is_tensor(lik) else torch.zeros(n_samples, device=device), float('-inf')) + if torch.is_tensor(lik): + lik = torch.where(invalid, neg_inf, lik).detach().clone() + else: + lik = torch.where(invalid, neg_inf, torch.zeros(n_samples, device=device)).detach().clone() + return Y, lik + else: + return Y diff --git a/tests/test_ms_selfcheck.py b/tests/test_ms_selfcheck.py new file mode 100644 index 0000000..c6c49c7 --- /dev/null +++ b/tests/test_ms_selfcheck.py @@ -0,0 +1,50 @@ +# tests/test_ms_selfcheck.py +import torch +from types import SimpleNamespace +from inc.samp_ms import samp_ms +from inc.diffus import SIR_STATES +import collections + +import numpy as np + +class DummyQNet: + def __init__(self, edge_index, T, device="cpu"): + self.T = int(T) + self.n_nodes = int(edge_index.max().item()) + 1 + self.device = torch.device(device) + neigh = [[] for _ in range(self.n_nodes)] + for u, v in edge_index.t().tolist(): + neigh[u].append(v); neigh[v].append(u) + self.neighbs = [torch.tensor(n, device=self.device, dtype=torch.long) for n in neigh] + deg = torch.tensor([len(n) for n in neigh], device=self.device).long() + self.rem = (deg + 1).view(-1, 1) # (nodes,1) + self.n_inf = torch.full((self.n_nodes, 1), self.n_nodes + 2, device=self.device) + self.zero = torch.tensor(0.0, device=self.device) + + def forward(self, yT, orig=True): + T, n, dev = self.T, self.n_nodes, self.device + zI = torch.zeros(T, n, 1, device=dev) + zR = torch.zeros(T, n, 1, device=dev) + return zI, zR, zI, zR + +def test_selfcheck(): + edge_index = torch.tensor([[0,1,1,2,2,3,3,4], + [1,0,2,1,3,2,4,3]], dtype=torch.long) + T, n = 6, 5 + q_net = DummyQNet(edge_index, T) + y3 = torch.tensor([0,1,1,0,0], dtype=torch.long) + y5 = torch.tensor([0,2,1,1,0], dtype=torch.long) + y_T = y5.clone() + obs_times = [3, 5, T] + obs_states = torch.stack([y3, y5, y_T]) + + zI0, zR0, zI, zR = q_net.forward(y_T, orig=True) + + Y, lq = samp_ms(q_net, y_T, zI, zR, n_samples=8, compute_lik=True, + obs_times=obs_times, obs_states=obs_states) + + assert Y.shape == (T, n, 8) + assert lq.shape == (8,) + assert (Y[1:] >= Y[:-1]).all().item() + assert torch.equal(Y[3, :, 0], y3) + assert torch.equal(Y[5, :, 0], y5) From bf6153ae67afcacfc187d943bce493325a435367 Mon Sep 17 00:00:00 2001 From: Yijing-Zuo Date: Fri, 26 Sep 2025 14:35:06 -0500 Subject: [PATCH 02/19] feat(ms): refine samp_ms, fix broadcasting & obs hard-constraints; add selfcheck test --- inc/samp_ms.py | 10 ++++++---- tests/test_ms_selfcheck.py | 1 - 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/inc/samp_ms.py b/inc/samp_ms.py index af6fbe5..4c9f9df 100644 --- a/inc/samp_ms.py +++ b/inc/samp_ms.py @@ -149,8 +149,9 @@ def samp_ms( invalid |= need_I # 覆盖本节点排序坐标 (t, i) 的采样与对数概率 xI[t, i] = SIR_STATES.S + qi = qI[t, i, 0] # qI[t, i] 形状为 (1,),扩展到 (samples,) - lI[t, i] = torch_log(qI[t, i].expand(n_samples)) + lI[t, i].fill_(float(torch_log(qi))) # 强制发生时仍需 opt 通过;否则该样本 invalid bad_no_opt = msk[u] & (~opt) if bad_no_opt.any(): @@ -161,7 +162,8 @@ def samp_ms( if need_I.any(): invalid |= need_I xI[t, i] = SIR_STATES.I - lI[t, i] = torch_log(1. - qI[t, i].expand(n_samples)) + qi = qI[t, i, 0] + lI[t, i].fill_(float(torch_log(1. - qi))) # 不发生时,无需 opt(与原版一致:opt=False 时不会进入 toss) else: # 目标为 R:本步(I->S 阶段)无法直接达到 R,不可达 @@ -170,9 +172,9 @@ def samp_ms( y[u] = torch.where(msk_opt, xI[t, i], y[u]) trs = (y[u] != SIR_STATES.I)# 是否发生了 I->S(= 正向新感染) - rem[u] = torch.where(msk[u], torch.where(trs, rem[u] - 1, q_net.n_inf), rem[u]) + rem[u] = torch.where(msk[u], torch.where(trs, rem[u] - 1, q_net.n_inf[u].expand_as(rem[u])), rem[u]) rem[vid] = torch.where(msk[u].unsqueeze(0), - torch.where(trs.unsqueeze(0), rem[vid] - 1, q_net.n_inf), + torch.where(trs.unsqueeze(0), rem[vid] - 1, q_net.n_inf[vid],expand_as(rem[vid])), rem[vid]) #最终 msk[u] 仅保留“进入 toss 且 opt 通过”的样本(与原版一致) msk[u] = msk_opt diff --git a/tests/test_ms_selfcheck.py b/tests/test_ms_selfcheck.py index c6c49c7..29553a1 100644 --- a/tests/test_ms_selfcheck.py +++ b/tests/test_ms_selfcheck.py @@ -2,7 +2,6 @@ import torch from types import SimpleNamespace from inc.samp_ms import samp_ms -from inc.diffus import SIR_STATES import collections import numpy as np From b7e33d6196ae9f85205d52404d14b5aa7306b6a7 Mon Sep 17 00:00:00 2001 From: Yijing-Zuo Date: Sat, 24 Jan 2026 06:25:19 -0600 Subject: [PATCH 03/19] WIP: migrate to new machine (code + input only) --- .gitignore | 12 + cri_ms.py | 292 ++++++++++++++ dhrec_ms.py | 224 +++++++++++ ditto_ms.py | 756 +++++++++++++++++++++++++++++++++++++ gcn_ms.py | 178 +++++++++ gin_ms.py | 156 ++++++++ inc/data.py | 3 +- inc/samp_ms.py | 194 ---------- inc/test_ms.py | 136 +++++++ input/covid/dist-covid.pt | Bin 0 -> 947883 bytes mkdir | 0 scripts/exp_sir_profile.py | 112 ++++++ tests/test_ms_selfcheck.py | 49 --- 13 files changed, 1867 insertions(+), 245 deletions(-) create mode 100644 cri_ms.py create mode 100644 dhrec_ms.py create mode 100644 ditto_ms.py create mode 100644 gcn_ms.py create mode 100644 gin_ms.py delete mode 100644 inc/samp_ms.py create mode 100644 inc/test_ms.py create mode 100644 input/covid/dist-covid.pt create mode 100644 mkdir create mode 100644 scripts/exp_sir_profile.py delete mode 100644 tests/test_ms_selfcheck.py diff --git a/.gitignore b/.gitignore index 5487d9d..a9794d5 100644 --- a/.gitignore +++ b/.gitignore @@ -130,3 +130,15 @@ dmypy.json # Pyre type checker .pyre/ + +# IDE +.idea/ + +# experiment artifacts +output/ +logs/ +log2/ + +# python cache +__pycache__/ +*.pyc diff --git a/cri_ms.py b/cri_ms.py new file mode 100644 index 0000000..9091023 --- /dev/null +++ b/cri_ms.py @@ -0,0 +1,292 @@ +from inc.diffus import * +from inc.test import * + +import argparse +import numpy as np +import networkx as nx + + +def get_args(): + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type=str, help='dataset name') + parser.add_argument('--seed', type=int, help='random seed') + parser.add_argument('--data_dir', type=str, help='dataset folder') + parser.add_argument('--output', type=str, help='output file name') + parser.add_argument('--device', type=torch.device, help='torch device') + + # Multi-snapshot options (aligned with inc/ditto_ms.py). + parser.add_argument( + '--obs_ts', type=str, default=None, + help='comma-separated observed time indices, e.g. "0,3,5"; ' + 'None means single-snapshot (use only the final snapshot)' + ) + parser.add_argument( + '--obs_k', type=int, default=None, + help='number of observed snapshots; if None, set to len(obs_ts) when obs_ts ' + 'is given, otherwise 1 (single-snapshot)' + ) + + args = parser.parse_args() + # normalize obs_ts / obs_k + if args.obs_ts is not None and len(args.obs_ts.strip()) > 0: + obs = [int(x) for x in args.obs_ts.split(',') if x.strip() != ''] + obs = sorted(set(obs)) + args.obs_ts = obs + if args.obs_k is None: + args.obs_k = len(obs) + else: + args.obs_ts = None + if args.obs_k is None: + args.obs_k = 1 + return args + + +def _bfs_dist(G: nx.Graph, s: int, cutoff: int = None) -> dict: + """ + Return {node: hop_distance} for nodes reachable from s. + + cutoff: optional BFS depth limit. When we only care whether dist <= T_thr, + setting cutoff=T_thr can be much faster on large graphs. + """ + if cutoff is None: + return dict(nx.single_source_shortest_path_length(G, s)) + return dict(nx.single_source_shortest_path_length(G, s, cutoff=int(cutoff))) + + +def cri_cluster(G: nx.Graph, obs_mask: np.ndarray, T_thr: int): + """ + CRI clustering (greedy k-center) — optimized for speed. + + Original bottleneck: + The legacy version repeatedly called BFS from *every infected node* inside + the greedy loop, causing huge runtimes on large VI. + + Fix (same concept, much faster): + Run BFS only from the current centers (with cutoff=T_thr), and maintain + each infected node's distance to its nearest center incrementally. + + Output: + clusters: list of infected-node lists (one per center) + VI: list of all infected nodes + """ + n = int(obs_mask.shape[0]) + VI = [u for u in range(n) if obs_mask[u] == 1] + if len(VI) == 0: + return [], VI + if len(VI) == 1: + return [VI], VI + + # If T_thr <= 0, every infected node must be its own center (no BFS needed). + if T_thr <= 0: + return [[u] for u in VI], VI + + VI_arr = np.asarray(VI, dtype=np.int64) + idx_of = {int(u): i for i, u in enumerate(VI_arr)} # node -> index in VI + + # --- 2-BFS "double sweep" heuristic to pick an initial far-apart pair --- + # We do NOT use cutoff here; only 2 BFS calls, so it's cheap and gives a better pair. + u0 = int(VI_arr[0]) + dist0 = _bfs_dist(G, u0, cutoff=None) + + u1 = next((int(u) for u in VI_arr if int(u) not in dist0), None) + if u1 is None: + u1 = max((int(u) for u in VI_arr), key=lambda u: dist0.get(u, -1)) + + dist1 = _bfs_dist(G, u1, cutoff=None) + u2 = next((int(u) for u in VI_arr if int(u) not in dist1), None) + if u2 is None: + u2 = max((int(u) for u in VI_arr), key=lambda u: dist1.get(u, -1)) + + centers = [u1] + if u2 != u1: + centers.append(u2) + + # --- Greedy k-center with incremental nearest-center distances --- + is_center = np.zeros(len(VI_arr), dtype=bool) + d_near = np.full(len(VI_arr), np.inf, dtype=np.float32) + + # For each center s, store only distances to infected nodes (shape: [|VI|]). + dist_to_VI = {} # center -> np.ndarray(|VI|,) + + def add_center(s: int): + nonlocal d_near + if s in dist_to_VI: + return + + # Key speedup: cutoff=T_thr (we only need to know if dist <= T_thr) + dist_s = _bfs_dist(G, s, cutoff=T_thr) + ds = np.fromiter((dist_s.get(int(u), np.inf) for u in VI_arr), + dtype=np.float32, count=len(VI_arr)) + dist_to_VI[s] = ds + + if s in idx_of: + is_center[idx_of[s]] = True + + d_near = np.minimum(d_near, ds) + d_near[is_center] = 0.0 + + for s in centers: + add_center(int(s)) + + with tqdm(desc='cluster', leave=False) as pbar: + while True: + # farthest infected node from current centers (excluding centers) + if (~is_center).any(): + tmp = d_near.copy() + tmp[is_center] = -1.0 + far_i = int(np.argmax(tmp)) + far_d = float(tmp[far_i]) + else: + far_i, far_d = -1, -1.0 + + if far_d <= float(T_thr): + break + + far_node = int(VI_arr[far_i]) + centers.append(far_node) + add_center(far_node) + pbar.update(1) + + # --- Assign each infected node to its nearest center (vectorized) --- + center_list = list(dict.fromkeys(centers)) # unique, keep order + dist_mat = np.vstack([dist_to_VI[s] for s in center_list]) # (k, |VI|) + assign = dist_mat.argmin(axis=0) # (|VI|,) + + clusters_map = {s: [] for s in center_list} + for i, u in enumerate(VI_arr): + s = center_list[int(assign[i])] + clusters_map[s].append(int(u)) + + clusters = [lst for lst in clusters_map.values() if len(lst) > 0] + return clusters, VI + + +def cri_rev_infect(G: nx.Graph, VI_all, Vi, y): + """ + Reverse infection for a cluster Vi (same as legacy version, but faster): + + - Expand BFS "wavefronts" tagged by sources x in Vi (pairs (u, x)). + - Stop once any node receives all tags. + - Choose the best center s (min sum of tag distances). + - For each x in Vi: set predicted infection time tI[x] = dist(s, x), + and mark y[x, tI[x]:] = 1. + + Speed fixes: + - Avoid allocating n empty dicts: create per-node dicts on-demand. + - Avoid O(n) scan each layer (track max_seen incrementally). + - Avoid final O(n) scan for candidates (track candidates during BFS). + """ + n = G.number_of_nodes() + ni = len(Vi) + if ni == 0: + return + + # g[u] is a dict {x: dist(u, x)}; allocate lazily + g = [None] * n + + def has_label(u: int, x: int) -> bool: + du = g[u] + return (du is not None) and (x in du) + + frontier = set() + for x in Vi: + x = int(x) + frontier.add((x, x)) + for v in G[x]: + frontier.add((int(v), x)) + + t_layer = 0 + max_seen = 0 + candidates = set() + + with tqdm(desc='rev_infect.expand', leave=False) as pbar: + while frontier and max_seen < ni: + next_frontier = set() + for u, x in frontier: + if has_label(u, x): + continue + if g[u] is None: + g[u] = {} + g[u][x] = t_layer + + lu = len(g[u]) + if lu > max_seen: + max_seen = lu + if lu == ni: + candidates.add(u) + + for v in G[u]: + v = int(v) + if not has_label(v, x): + next_frontier.add((v, x)) + + frontier = next_frontier + t_layer += 1 + pbar.update(1) + + if not candidates: + # Fallback (should be rare if clustering worked): + x0 = int(Vi[0]) + y[x0, 0:] = 1 + return + + s = min(candidates, key=lambda u: sum(g[u].values())) + + # Set first-infection time for each tagged x in the cluster and mark trajectory + tI_map = g[s] + for x, t0 in tI_map.items(): + t0 = int(t0) + if 0 <= x < y.shape[0]: + t0 = max(0, min(t0, y.shape[1] - 1)) # clamp to [0, T] + y[x, t0:] = 1 + + +@torch.no_grad() +def cri_ms_run(data): + """ + Multi-snapshot CRI: + - For each observed time t_obs: + * extract infected mask (I=1, others=0), + * cluster with radius threshold = t_obs, + * reverse-infect per cluster to get first-infection times, + * mark y[:, tI: ] = 1 for those nodes, + - Merge across snapshots by OR (sum since we fill with 1's from tI onward). + """ + T = int(data.T.item()) + n_nodes = int(data.num_nodes) + n_cls = int(data.y.max().item() + 1) + + # Determine observed times to use + if args.obs_ts is not None and len(args.obs_ts) > 0: + obs_times = [t for t in args.obs_ts if 0 <= t <= T] + obs_times = sorted(set(obs_times)) + if len(obs_times) == 0: + obs_times = [T] + else: + obs_times = [T] + + # Build graph + G = pyg.utils.to_networkx(data, to_undirected=True, remove_self_loops=True) + + # Accumulator for predictions + y_pred = np.zeros((n_nodes, T + 1), dtype=np.int32) + + for t_obs in tqdm(obs_times, desc='obs_times', leave=False): + obs_vec = (data.y[:, t_obs].cpu().detach().numpy() & 1).astype(np.int32) + if obs_vec.sum() == 0: + continue + + clusters, VI_all = cri_cluster(G, obs_vec, T_thr=int(t_obs)) + for Vi in tqdm(clusters, desc=f'rev_infect@t={t_obs}', leave=False): + cri_rev_infect(G, VI_all, Vi, y_pred) + + return torch.tensor(np.minimum(y_pred, n_cls - 1), dtype=torch.long, device=data.y.device) + + +# ---- entry point ---- +args = get_args() +seed_all(args.seed) +tester = Tester(args.data_dir, args.device, cri_ms_run) +tester.test([args.dataset], rep=1) +tester.save(args.output) + diff --git a/dhrec_ms.py b/dhrec_ms.py new file mode 100644 index 0000000..dd73648 --- /dev/null +++ b/dhrec_ms.py @@ -0,0 +1,224 @@ +from inc.diffus import * +from inc.test import * + + +def get_args(): + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type=str, help='dataset name') + parser.add_argument('--seed', type=int, help='random seed') + parser.add_argument('--data_dir', type=str, help='dataset folder') + parser.add_argument('--output', type=str, help='output file name') + parser.add_argument('--device', type=torch.device, help='torch device') + + # diffusion parameter estimation (kept as in the baseline) + parser.add_argument('--b_pI0', type=float, help='initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type=float, help='initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type=int, help='optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type=float, help='learning rate in diffusion parameter estimation') + + # multi-snapshot options (aligned with ditto_ms.py) + parser.add_argument( + '--obs_ts', type=str, default=None, + help='comma-separated observed time indices, e.g. "0,3,5"; ' + 'None means single-snapshot (use only the final snapshot)' + ) + parser.add_argument( + '--obs_k', type=int, default=None, + help='number of observed snapshots; if None, set to len(obs_ts) when obs_ts is given, ' + 'otherwise 1 (single-snapshot)' + ) + args = parser.parse_args() + + # parse obs_ts in the same spirit as inc/ditto_ms.py + if args.obs_ts is not None and len(args.obs_ts.strip()) > 0: + obs = [int(x) for x in args.obs_ts.split(',') if x.strip() != ''] + obs = sorted(set(obs)) + args.obs_ts = obs + if args.obs_k is None: + args.obs_k = len(obs) + else: + args.obs_ts = None + if args.obs_k is None: + args.obs_k = 1 + return args + + +def pcdsvc_greedy(bpar, G, y): + """ + One-step backward inference (t -> t-1) under the PCDSVC-style greedy rule. + This is the original single-snapshot kernel kept unchanged; it maps a snapshot y(t) + to a previous snapshot x(t-1). + + Fix: + - Some SI datasets legitimately have pR = 0 (no recovery). In that case, lr = -log(pR) + should NOT be evaluated (it is unused anyway because there is no R state in y). + - Also guard against boundary numeric issues (pI -> 1, pR -> 0) to avoid log(0). + """ + n = G.number_of_nodes() + + # ---- numeric guards / SI-safe handling ---- + eps = 1e-12 + + # pI is used in all cases; only clip the upper bound to avoid log(0) at (1 - pI). + pI = float(getattr(bpar, 'pI', 0.0)) + if not np.isfinite(pI): + pI = 0.0 + pI = min(max(pI, 0.0), 1.0 - eps) + l1s = -np.log(1.0 - pI) + + # pR is only needed when there are recovered nodes in the current snapshot. + # For SI datasets, y never contains state=2, so skip log(pR) entirely. + if np.any(y == 2): + pR = float(getattr(bpar, 'pR', 0.0)) + if not np.isfinite(pR): + pR = eps + pR = min(max(pR, eps), 1.0) + lr = -np.log(pR) + else: + lr = 0.0 + # ------------------------------------------ + + # x: 2 means "unknown yet / start from R" in the original code, then moves to 1 or 0 + x = np.where(y == 2, 2, 0) + + we = l1s + ws = np.zeros(n, dtype=np.float32) + wi = np.zeros(n, dtype=np.float32) + + for u in range(n): + if y[u] == 0: # S + for v in G.neighbors(u): + wi[v] += l1s + elif y[u] == 1: # I + ws[u] -= 1. + else: # R + ws[u] += lr - 1. + wi[u] += lr + + # R --> I + pbar = tqdm(disable=True) + while True: + mvs = [] + for u in range(n): + if y[u] >= 1 and x[u] != 1: + cur = wi[u] + for v in G.neighbors(u): + if y[v] >= 1 and x[v] != 1: + cur -= we + mvs.append((cur, u)) + if len(mvs) == 0: + break + mv = min(mvs, key=lambda mv: mv[0]) + if mv[0] >= 0.: + dom = True + for u in range(n): + if y[u] == 1: + dm = (x[u] == 1) + for v in G.neighbors(u): + dm |= (x[v] == 1) + if dm: + break + dom &= dm + if dom: + break + x[mv[1]] = 1 + pbar.update(1) + + # I --> S + for u in range(n): + if x[u] == 2 and ws[u] < 0: + dm = False + for v in G.neighbors(u): + dm |= (x[v] == 1) + if dm: + break + if dm: + x[u] = 0 + pbar.update(1) + pbar.close() + return x + + +def _parse_obs_ts(args, T): + """ + Build a sorted & unique list of observed time indices within [0, T], + ensuring the final snapshot T is always included as the anchor. + """ + if args.obs_ts is None: + obs_ts = [] + else: + obs_ts = [t for t in args.obs_ts if 0 <= t <= T] + + if T not in obs_ts: + obs_ts.append(T) + obs_ts = sorted(set(obs_ts)) + return obs_ts + + +def _build_observation_map(data, obs_ts): + """ + Return {t: np.array states at time t} for all observed times t. + """ + y = data.y # (nodes, T+1) + obs = {} + for t in obs_ts: + obs[t] = y[:, t].cpu().detach().numpy() + return obs + + +def pcdsvc_run(data): + """ + Multi-snapshot backward reconstruction: + - Split the timeline by observed time points (including T as anchor). + - For each segment [t_prev, t_next], start from the observed y(t_next) + and repeatedly apply the single-step greedy kernel to obtain y(t_prev+1),...,y(t_prev), + clamping to ground-truth observation whenever we hit an observed time. + """ + bpar = b_estim(data, args) # estimate diffusion parameters once + + with torch.no_grad(): + T = int(data.T.item()) + n_nodes = data.num_nodes + n_cls = int(data.y.max().item() + 1) + + # build graph & observations + G = pyg.utils.to_networkx(data, to_undirected=True, remove_self_loops=True) + obs_ts = _parse_obs_ts(args, T) # ensure T is present + obs_map = _build_observation_map(data, obs_ts) + + # init buffer and seed observed frames + y_pred = np.zeros((n_nodes, T + 1), dtype=np.int32) + for t in obs_ts: + y_pred[:, t] = obs_map[t] + + # include 0 to cover the entire span; we will reconstruct down to t=0 + time_cuts = sorted(set(obs_ts + [0])) + + # process segments from later to earlier + for idx in trange(len(time_cuts) - 1, 0, -1, desc='ms-backtrack'): + t_prev = time_cuts[idx - 1] + t_next = time_cuts[idx] # t_next > t_prev + y_cur = y_pred[:, t_next] # this is observed or already clamped + + # step-by-step: t_next -> t_prev + for t in range(t_next, t_prev, -1): + y_prev = pcdsvc_greedy(bpar, G, y_cur) + + # if the new time (t-1) is observed, clamp to the observation + if (t - 1) in obs_map: + y_prev = obs_map[t - 1] + + y_pred[:, t - 1] = y_prev + y_cur = y_prev + + # clip to valid classes and return torch tensor on the original device + return torch.tensor(np.minimum(y_pred, n_cls - 1), + dtype=torch.long, device=data.y.device) + + +args = get_args() +seed_all(args.seed) +tester = Tester(args.data_dir, args.device, pcdsvc_run) +tester.test([args.dataset], rep=1) +tester.save(args.output) + diff --git a/ditto_ms.py b/ditto_ms.py new file mode 100644 index 0000000..055b3b5 --- /dev/null +++ b/ditto_ms.py @@ -0,0 +1,756 @@ +from inc.diffus import * +from inc.nn import * +from inc.test import * + +def get_args(): + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type = str, help = 'dataset name') + parser.add_argument('--seed', type = int, help = 'random seed') + parser.add_argument('--data_dir', type = str, help = 'dataset folder') + parser.add_argument('--output', type = str, help = 'output file name') + parser.add_argument('--device', type = torch.device, help = 'torch device') + parser.add_argument('--b_pI0', type = float, help = 'initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type = float, help = 'initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type = int, help = 'optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type = float, help = 'learning rate in diffusion parameter estimation') + parser.add_argument('--q_steps', type = int, help = 'training steps for the proposal model') + parser.add_argument('--q_lr', type = float, help = 'learning rate for the proposal model') + parser.add_argument('--q_hid', type = int, help = 'hidden size of the proposal model') + parser.add_argument('--q_gnn', type = int, help = 'number of layers of the GNN in the proposal model') + parser.add_argument('--q_mlp', type = int, help = 'number of layers of the MLP in the proposal model') + parser.add_argument('--q_samples', type = int, help = 'sample size to estimate the loss function of the proposal model') + parser.add_argument('--q_zlim', type = int, help = 'a hyperparameter to stablize gradient') + parser.add_argument('--p_coef', type = float, help = 'the coefficient gamma in the initial distribution P[y_0]') + parser.add_argument('--t_samples', type = int, help = 'MCMC sample size') + parser.add_argument('--t_steps', type = int, help = 'MCMC steps') + parser.add_argument('--t_keep', type = float, help = 'moving average in MCMC') + parser.add_argument('--obs_time', type = str, default = '', help = 'extra observed snapshot times, comma-separated, e.g., 5,7,9') + # Multi-snapshot proposal diagnostics/safety. + # The backward segment sampler uses rejection sampling under hard constraints. + # If a segment is infeasible (empty support) or the proposal assigns vanishing + # probability to feasible states, the rejection loop can otherwise run forever. + parser.add_argument('--ms_max_rounds', type=int, default=2048, + help='max rejection rounds per backward step in multi-snapshot sampling') + args = parser.parse_args() + return args + +class QNet(nn.Module): + @classmethod + def make(cls, data, args): + return cls( + eidx = data.edge_index, + T = data.T.item(), + hid = args.q_hid, + gnn = args.q_gnn, + mlp = args.q_mlp, + n_nodes = data.num_nodes, + zlim = args.q_zlim, + ms_max_rounds = getattr(args, 'ms_max_rounds', 2048), + ).to(args.device) + def __init__(self, eidx, T, hid, gnn, mlp, n_nodes, zlim, ms_max_rounds=2048): + super().__init__() + self.eidx = eidx + self.device = self.eidx.device + self.n_nodes = n_nodes + self.n_inf = self.n_nodes + 2 + self.n_edges = self.eidx.size(dim = 1) + self.zlim = zlim + # Maximum rejection rounds per backward step in multi-snapshot segment sampling. + # This avoids infinite loops when a segment has empty/tiny support under hard constraints. + self.ms_max_rounds = int(ms_max_rounds) + self.T = T + self.hid = int(hid) + self.gnn_dep = int(gnn) + self.mlp_dep = int(mlp) + self.w = nn.Parameter(data = torch.randn((self.n_edges, self.hid), dtype = torch.float32, device = self.device), requires_grad = True) + self.gnn = GNN(v_in = 1, e_in = self.hid, hid = self.hid, dep = self.gnn_dep) + self.mlp = MLP([self.hid] * self.mlp_dep + [2 * self.T]) + self.rem = (pyg.utils.degree(self.eidx[1], num_nodes = self.n_nodes).long().unsqueeze(dim = 1) + 1).detach().clone() # (nodes, 1) + self.neighbs = [[] for u in range(self.n_nodes)] + for i in range(self.n_edges): + self.neighbs[self.eidx[0, i].item()].append(self.eidx[1, i].item()) + for u in range(self.n_nodes): + self.neighbs[u] = torch.tensor(self.neighbs[u], dtype = torch.long, device = self.device) + self.adj = torch.sparse_coo_tensor( + indices = torch.stack([self.eidx[1], self.eidx[0]], dim = 0), + values = torch.ones(self.n_edges, dtype = torch.float32, device = self.device), + size = (self.n_nodes, self.n_nodes), + ).coalesce() + self.zero = torch.tensor(0., dtype = torch.float, device = self.device) + def clamp_z(self, z): + return z.clamp(-self.zlim, self.zlim) + def forward(self, y, orig = False): # y: (nodes, samples) + n_nodes, n_samples = y.size() + y = y.T.reshape((-1, 1)) # (samples*nodes, 1) + eidx = (self.eidx.unsqueeze(dim = 1) + n_nodes * torch.arange(n_samples, dtype = torch.long, device = y.device).unsqueeze(dim = -1)).reshape((2, -1)) # (2, samples*edges) + w = self.w.repeat(n_samples, 1) # (samples, hid) + z, e = self.gnn(y.float(), eidx, w) + z = self.mlp(z) # (samples*nodes, 2*T) + z = z.T.reshape((2 * self.T, n_samples, -1)) # (2*T, samples, nodes) + zI, zR = z[: self.T], z[self.T :] # (T, samples, nodes) + zI, zR = zI.transpose(1, 2), zR.transpose(1, 2) # (T, nodes, samples) + if orig: + return zI, zR, self.clamp_z(zI), self.clamp_z(zR) + else: + return self.clamp_z(zI), self.clamp_z(zR) + def lik(self, Y): # Y: (T+1, nodes, samples) + n_samples = Y.size(dim = 2) + zI0, zR0, zI, zR = self.forward(Y[-1], orig = True) # (T, nodes, samples) + zI = zI.clone().detach().requires_grad_(True); zI.retain_grad() + zR = zR.clone().detach().requires_grad_(True); zR.retain_grad() + # R->I + qR = torch.sigmoid(zR) # (T, nodes, samples) # prob of R->I + lR1 = torch_log(qR) # (T, nodes, samples) + lR0 = torch_log(1. - qR) # (T, nodes, samples) + with torch.no_grad(): + mskR = (Y[1 :] == SIR_STATES.R) # (T, nodes, samples) + trsR = (Y[: -1] != SIR_STATES.R) # (T, nodes, samples) + # I->S + zI_, uid = zI.sort(dim = 1, descending = True) # (T, nodes, samples) + qI = torch.sigmoid(zI_) # (T, nodes, samples) # prob of I->S + lI1 = torch_log(qI) # (T, nodes, samples) + lI0 = torch_log(1. - qI) # (T, nodes, samples) + with torch.no_grad(): + mskI = ((Y[1 :] >= SIR_STATES.I) & (Y[: -1] <= SIR_STATES.I)).flatten() # (T * nodes * samples) + trsI = ((Y[: -1] != SIR_STATES.I)).flatten() # (T * nodes * samples) + rem = torch.where(mskI, self.rem.expand(self.T, -1, n_samples).flatten(), self.n_inf) # (T * nodes * samples) + ptr = torch.arange(self.T, dtype = torch.long, device = self.device).unsqueeze(dim = 1) * self.n_nodes # (T, 1) + for i in range(uid.size(dim = 1)): + uidi = (ptr + uid[:, i]).flatten() * n_samples # (T * samples) + mski = mskI[uidi] # (T * samples) + if mski.max(): + trsi = trsI[uidi] # (T * samples) + vids, degi = [], [0] + for t in range(self.T): + for j in range(n_samples): + u = uid[t, i, j] + vids.append((t * self.n_nodes + u.unsqueeze(dim = 0)) * n_samples) + vid = self.neighbs[u.item()] + vids.append((t * self.n_nodes + vid) * n_samples) + degi.append(vid.size(dim = 0) + 1) + vids = torch.cat(vids, dim = 0) # (T * sum neighbs) + degi = torch.tensor(degi, dtype = torch.long, device = self.device) # (1 + T * samples) + indptr = degi.cumsum(dim = 0) # (1 + T * samples) + degi = degi[1 :] # (T * samples) + rems = rem.flatten()[vids] # (T * sum neighbs) + opti = (pysc.segment_min_csr(src = rems, indptr = indptr)[0] > 1) # (T * samples) + rem.flatten()[vids] = torch.where(mski.repeat_interleave(repeats = degi), torch.where(trsi.repeat_interleave(repeats = degi), rems - 1, self.n_inf), rems) # (T * sum neighbs) + mskI[uidi] &= opti # (T * samples) + # likR + likI + lik = ( + torch.where(mskR, torch.where(trsR, lR1, lR0), self.zero).reshape(-1, n_samples) + + torch.where( + mskI.reshape(-1, n_samples), + torch.where( + trsI.reshape(-1, n_samples), + lI1.reshape(-1, n_samples), + lI0.reshape(-1, n_samples), + ), + self.zero, + ) + ).sum(dim = 0) # (samples,) + return lik, zI0, zR0, zI, zR # (samples,) + @torch.no_grad() + def clamp_grad(self, z0, grad): + return torch.where(z0 < self.zlim, torch.where(z0 > -self.zlim, grad, F.relu(grad)), -F.relu(-grad)) + def backward(self, loss, zI0, zR0, zI, zR): + loss.backward() + z0 = torch.stack([zI0, zR0], dim = 0) + z0.backward(torch.stack([self.clamp_grad(zI0, zI.grad), self.clamp_grad(zR0, zR.grad)], dim = 0)) + @torch.no_grad() + def samp(self, y, zI, zR, n_samples, compute_lik = False): # y: (nodes,); zI, zR: (T, nodes, 1) + zI, uid = zI.sort(dim = 1, descending = True) # (T, nodes, 1) + uid = uid.squeeze(dim = 2) # (T, nodes) + qI = torch.sigmoid(zI) # (T, nodes, 1) # prob of I->S + xI = SIR_STATES.I - qI.expand(-1, -1, n_samples).bernoulli().long() # (T, nodes, samples) # 1 for I->S + lI = torch_log(torch.where(xI != SIR_STATES.I, qI, 1. - qI)) # (T, nodes, samples) + qR = torch.sigmoid(zR) # (T, nodes, 1) # prob of R->I + xR = SIR_STATES.R - qR.expand(-1, -1, n_samples).bernoulli().long() # (T, nodes, samples) # 1 for R->I + lR = torch_log(torch.where(xR != SIR_STATES.R, qR, 1. - qR)) # (T, nodes, samples) + y = y.unsqueeze(dim = 1).expand(-1, n_samples) # (nodes, samples) + Y = torch.empty(self.T, self.n_nodes, n_samples, dtype = torch.long, device = self.device) # (T, nodes, samples) + if compute_lik: + lik = self.zero + for t in range(self.T - 1, -1, -1): + # R->I + msk = (y == SIR_STATES.R) # (nodes, samples) + y = torch.where(msk, xR[t], y) # (nodes, samples) + if compute_lik: + lik = lik + torch.where(msk, lR[t], self.zero).sum(dim = 0) # (samples,) + # I->S + msk = (y == SIR_STATES.I) # (nodes, samples) + rem = torch.where(msk, self.rem, self.n_inf) # (nodes, samples) + for i, u in enumerate(uid[t]): + if msk[u].max(): + vid = self.neighbs[u.item()] # (neighbs,) + opt = (rem[u] > 1) & (rem[vid].min(dim = 0).values > 1) # (samples,) + msk_opt = msk[u] & opt + y[u] = torch.where(msk_opt, xI[t, i], y[u]) # (samples,) + trs = (y[u] != SIR_STATES.I) # (samples,) + rem[u] = torch.where(msk[u], torch.where(trs, rem[u] - 1, self.n_inf), rem[u]) # (samples,) + rem[vid] = torch.where(msk[u].unsqueeze(dim = 0), torch.where(trs.unsqueeze(dim = 0), rem[vid] - 1, self.n_inf), rem[vid]) # (neighbs, samples) + msk[u] = msk_opt + Y[t] = y + if compute_lik: + lik = lik + torch.where(msk[uid[t]], lI[t], self.zero).sum(dim = 0) # (samples,) + Y = Y.detach().clone() + if compute_lik: + lik = lik.detach().clone() + return Y, lik + else: + return Y + + @torch.no_grad() + def ext_ok(self, x, yL, t, TL): + d = t - TL + ok = (x >= yL.unsqueeze(dim=1)).all(dim=0) # (samples,) + if not ok.max(): + return ok + src = (yL == SIR_STATES.I).unsqueeze(dim=1) # (nodes, 1) + A = (x != SIR_STATES.S) & (yL != SIR_STATES.R).unsqueeze(dim=1) # (nodes, samples) + reach = src.expand(-1, x.size(dim=1)) & A # (nodes, samples) + for _ in range(d): + nbr = (torch.sparse.mm(self.adj, reach.float()) > 0) & A # (nodes, samples) + reach = reach | nbr + Iset = (x == SIR_STATES.I) & (yL != SIR_STATES.R).unsqueeze(dim=1) # (nodes, samples) + Rset = (x == SIR_STATES.R) & (yL != SIR_STATES.R).unsqueeze(dim=1) # (nodes, samples) + ok = ok & ~(Iset & ~reach).any(dim=0) + # NOTE: allow one-step S->R (infection + recovery within the same discrete step), + # so R-nodes at time t only need distance <= d (not d-1). + ok = ok & ~(Rset & ~reach).any(dim=0) + return ok + + @torch.no_grad() + def _samp_step(self, y, zI, uid, zR, t, compute_lik=False, yL=None): + """One backward step: sample y_t given y_{t+1}=y. + + This follows DITTO's original local-support (right-end feasibility) design via the + `msk/rem` mechanism. + + Multi-snapshot add-on (Fix 2): if a left endpoint snapshot yL (= y_{TL}) is provided, + we *hard clamp* the **left-monotonicity** necessary constraint directly inside the + sampler so we never generate y_t < yL. + + Concretely (S < I < R): + - Nodes with yL==R must stay R for all t>=TL => disable backward R->I. + - Nodes with yL==I must stay in {I,R} => disable backward I->S on those nodes. + + IMPORTANT: We enforce this by masking *sampling choices*, not by post-hoc overwriting, + so the returned `lik` remains the correct proposal log-probability. + + Parameters + ---------- + y : LongTensor, (nodes, samples) + Current snapshot y_{t+1} for a batch of samples. + yL : LongTensor or None, (nodes,) + Left observed snapshot y_{TL} (for monotonic clamping). If None, no clamping. + """ + n_samples = y.size(dim=1) + + # Left-monotonic hard constraints (segment-wise), if provided. + if yL is not None: + force_R = (yL == SIR_STATES.R).unsqueeze(dim=1) # (nodes, 1) + force_I = (yL == SIR_STATES.I) # (nodes,) + else: + force_R = None + force_I = None + if compute_lik: + lik = self.zero + # R->I + qR = torch.sigmoid(zR[t]) # (nodes, 1) + xR = SIR_STATES.R - qR.expand(-1, n_samples).bernoulli().long() # (nodes, samples) + if compute_lik: + lR = torch_log(torch.where(xR != SIR_STATES.R, qR, 1. - qR)) # (nodes, samples) + msk = (y == SIR_STATES.R) # (nodes, samples) + # Fix 2 (part 1): nodes already recovered at TL must remain R => do NOT sample R->I. + if force_R is not None: + msk = msk & (~force_R) + y = torch.where(msk, xR, y) + if compute_lik: + lik = lik + torch.where(msk, lR, self.zero).sum(dim=0) # (samples,) + # I->S + qI = torch.sigmoid(zI[t]) # (nodes, 1) (already sorted by zI) + xI = SIR_STATES.I - qI.expand(-1, n_samples).bernoulli().long() # (nodes, samples) + if compute_lik: + lI = torch_log(torch.where(xI != SIR_STATES.I, qI, 1. - qI)) # (nodes, samples) + msk = (y == SIR_STATES.I) # (nodes, samples) + rem = torch.where(msk, self.rem, self.n_inf) # (nodes, samples) + for i, u in enumerate(uid[t]): + if msk[u].max(): + vid = self.neighbs[u.item()] # (neighbs,) + opt = (rem[u] > 1) & (rem[vid].min(dim=0).values > 1) # (samples,) + # Fix 2 (part 2): nodes infected at TL must never go below I => do NOT sample I->S. + # Enforce by forcing `opt=False` (deterministic keep-I) for those nodes. + if force_I is not None: + opt = opt & (~force_I[u]) + msk_opt = msk[u] & opt + y[u] = torch.where(msk_opt, xI[i], y[u]) # (samples,) + trs = (y[u] != SIR_STATES.I) # (samples,) + rem[u] = torch.where(msk[u], torch.where(trs, rem[u] - 1, self.n_inf), rem[u]) + rem[vid] = torch.where(msk[u].unsqueeze(dim=0), + torch.where(trs.unsqueeze(dim=0), rem[vid] - 1, self.n_inf), rem[vid]) + msk[u] = msk_opt + if compute_lik: + lik = lik + torch.where(msk[uid[t]], lI, self.zero).sum(dim=0) # (samples,) + return y, lik + else: + return y + + @torch.no_grad() + def _samp_seg(self, yR, zI, uid, zR, n_samples, TL, TR, yL=None, compute_lik=False): + """ + Sample ONE segment (TL, TR] in *reverse* temporal order (backward sampling). + + Segment definition: + - Left endpoint time TL (may be observed / clamped if yL is provided) + - Right endpoint time TR (always observed here; yR is the snapshot at TR) + - We generate snapshots for times: TR-1, TR-2, ..., TL (if yL is None) or TL+1 (if yL is fixed) + + Why segment-wise? + In multi-snapshot DITTO, we must satisfy *all* observed snapshots exactly. + We sample each segment backward from the fixed right endpoint y_TR. However, in the multi-snapshot + setting we also need a *left feasibility* guarantee: sampled states must still be extendable to + match the left observed snapshot y_TL. This is the left-extendability hard constraint Ext(t). + + Parameters + ---------- + yR : LongTensor, shape (nodes,) + Fixed right endpoint snapshot y_{TR}. + zI : FloatTensor, shape (T, nodes, 1) + Proposal logits (already sorted along nodes dim) controlling backward I->S decisions. + NOTE: must be consistent with uid. + uid : LongTensor, shape (T, nodes) + Node indices sorted by descending zI at each time t (DITTO's ordering trick). + zR : FloatTensor, shape (T, nodes, 1) + Proposal logits controlling backward R->I decisions (NOT sorted; original node order). + n_samples : int + How many independent histories to sample in parallel. + TL, TR : int + Segment endpoints (TL < TR). + yL : LongTensor or None + If provided, this is the *observed* left endpoint snapshot y_{TL} (hard constraint). + In that case we will NOT sample time TL; we clamp it to yL. + compute_lik : bool + If True, also return log-probability under the proposal (needed by M-H acceptance). + + Returns + ------- + Y_seg : LongTensor, shape (TR-TL, nodes, samples) + The sampled segment snapshots in *forward* time indexing within the segment: + Y_seg[k] corresponds to time (TL + k), for k=0..(TR-TL-1). + If yL is provided, then Y_seg[0] == yL is included (clamped). + lik_seg : FloatTensor, shape (samples,) (only if compute_lik=True) + Sum of log-probabilities of all backward steps performed inside this segment. + """ + L = TR - TL # number of time indices in [TL, TR) that we store in Y_seg + Y = torch.empty(L, self.n_nodes, n_samples, dtype=torch.long, device=self.device) + + if compute_lik: + # segment log-likelihood under the proposal Q_theta + lik = torch.zeros(n_samples, dtype=torch.float, device=self.device) + + # Current "right" snapshot y_{t+1}. Start from fixed right endpoint y_{TR}. + y = yR.unsqueeze(dim=1).expand(-1, n_samples) # (nodes, samples) + + # --------------------------------------------------------------------- + # Precompute yL-dependent tensors once per segment. + # --------------------------------------------------------------------- + if yL is not None: + yL_col = yL.unsqueeze(dim=1) # (nodes, 1) for monotonicity check x >= y_TL + yL_not_R = (yL != SIR_STATES.R).unsqueeze(dim=1) # (nodes, 1) exclude nodes fixed to R at TL + src = (yL == SIR_STATES.I).unsqueeze(dim=1) # (nodes, 1) infection sources at TL + + def ext_ok_fast(x, t): + """ + Faster ext_ok that reuses yL-related precomputations. + x: (nodes, samples) candidate snapshot at time t + """ + d = t - TL + ok = (x >= yL_col).all(dim=0) # (samples,) + if not ok.max(): + return ok + A = (x != SIR_STATES.S) & yL_not_R # (nodes, samples) + reach = src.expand(-1, x.size(dim=1)) & A # (nodes, samples) + for _ in range(d): + nbr = (torch.sparse.mm(self.adj, reach.float()) > 0) & A + new_reach = reach | nbr + if torch.equal(new_reach, reach): + reach = new_reach + break + reach = new_reach + Iset = (x == SIR_STATES.I) & yL_not_R + Rset = (x == SIR_STATES.R) & yL_not_R + ok = ok & ~(Iset & ~reach).any(dim=0) + ok = ok & ~(Rset & ~reach).any(dim=0) + return ok + else: + ext_ok_fast = None + + # --------------------------------------------------------------------- + # Segment-level empty-support check: the observed right endpoint itself + # must be extendable from the observed left endpoint (when present). + # --------------------------------------------------------------------- + if ext_ok_fast is not None: + okR = ext_ok_fast(yR.unsqueeze(dim=1), TR) # (1,) + if not bool(okR.item()): + n_src = int((yL == SIR_STATES.I).sum().item()) + n_fixR = int((yL == SIR_STATES.R).sum().item()) + n_yR_I = int((yR == SIR_STATES.I).sum().item()) + n_yR_R = int((yR == SIR_STATES.R).sum().item()) + mono = bool((yR >= yL).all().item()) + raise RuntimeError( + "[DITTO-MS] Segment infeasible (empty support) under hard constraints. " + f"Segment (TL={TL}, TR={TR}, len={TR-TL}). " + f"Monotonic(y_TR>=y_TL)={mono}. " + f"#src_I@TL={n_src}, #fixed_R@TL={n_fixR}, #I@TR={n_yR_I}, #R@TR={n_yR_R}. " + "This typically means the observations cannot be bridged on the graph within the time budget " + "(e.g., src is empty/small, isolated targets, or timestamps over-constrain the diffusion)." + ) + + # --------------------------------------------------------------------- + # First segment (TL=0) has no left-extendability constraint; keep original sampler. + # For subsequent segments, use constructive hop-layer growth to enforce Ext(t) by construction. + # --------------------------------------------------------------------- + if yL is None: + # no left constraint, sample purely by original backward local-support steps + for k in range(L - 1, -1, -1): + t = TL + k + if compute_lik: + y, lik_t = self._samp_step(y, zI, uid, zR, t, compute_lik=True, yL=None) + lik = lik + lik_t + else: + y = self._samp_step(y, zI, uid, zR, t, compute_lik=False, yL=None) + Y[k] = y + if compute_lik: + return Y.detach().clone(), lik.detach().clone() + else: + return Y.detach().clone() + + # --------------------------------------------------------------------- + # Constructive extendable sampler (Scheme A1): + # At each time t in (TL, TR): + # - pool := nodes with y_{t+1} != S and yL != R + # - build A_t (non-S set) by hop layers from src within pool: + # include all nodes within distance <= d-1 + # optionally include some nodes at exact distance d + # force-include distance-d nodes that are needed to infect distance-(d+1) nodes + # - set y_t outside A_t to S (or R if blocked) + # - inside A_t, decide I/R using zR, but force I on nodes needed as infection sources + # This eliminates the rejection loop for Ext(t). + # --------------------------------------------------------------------- + + # Build unsorted zI (needed to derive p_add on boundary layer in original node order). + # zI is sorted along nodes dim with permutation uid. + zI_sorted_ = zI.squeeze(dim=2) # (T, nodes) + zI_unsorted = torch.empty_like(zI_sorted_) # (T, nodes) + for tt in range(self.T): + zI_unsorted[tt, uid[tt]] = zI_sorted_[tt] + + # Precompute fixed masks from left observation. + blocked = (yL == SIR_STATES.R).unsqueeze(dim=1) # (nodes, 1) + yL_not_R = ~blocked + src = (yL == SIR_STATES.I).unsqueeze(dim=1) # (nodes, 1) + + # We do NOT sample time TL itself; clamp it to yL. + k_min = 1 + + for k in range(L - 1, k_min - 1, -1): + t = TL + k + d = t - TL # hop budget for Ext(t) + + y_next = y # y_{t+1}, shape (nodes, samples) + + # pool = {u: y_{t+1,u} != S} \ blocked + pool = (y_next != SIR_STATES.S) & yL_not_R # (nodes, samples) + + # If any sample has non-empty pool but no src in pool, segment is infeasible for that sample. + src_in_pool = (src.expand(-1, n_samples) & pool).any(dim=0) # (samples,) + if (~src_in_pool & pool.any(dim=0)).any(): + raise RuntimeError( + "[DITTO-MS] Constructive sampler hit infeasible intermediate state: " + f"at time t={t} (TL={TL},TR={TR}), pool non-empty but src not in pool for some samples. " + "This indicates either inconsistent observations or a bug in monotonic clamping." + ) + + # ----------------------------------------------------------------- + # Hop-layer BFS within pool from src, up to d+1 layers. + # We need: + # - reach_{d-1}: nodes within dist <= d-1 (must be non-S at time t to support outer growth) + # - layer_d: nodes at exact dist d (optional, but some are forced) + # - layer_{d+1}: nodes at exact dist d+1 (cannot be non-S at time t; must be new at t+1) + # ----------------------------------------------------------------- + frontier = src.expand(-1, n_samples) & pool # layer 0 + reach = frontier.clone() + reach_dminus1 = frontier.clone() # will be overwritten if d-1 >= 1 + layer_d = torch.zeros_like(frontier) + layer_d1 = torch.zeros_like(frontier) + + # Special: if d-1 == 0, then reach_dminus1 is just layer0. + # We'll record reach after step (d-1) as reach_dminus1. + for h in range(1, d + 2): # compute layers 1..d+1 + nbr = (torch.sparse.mm(self.adj, frontier.float()) > 0) & pool & (~reach) + frontier = nbr + reach = reach | frontier + if h == d - 1: + reach_dminus1 = reach.clone() + if h == d: + layer_d = frontier.clone() + if h == d + 1: + layer_d1 = frontier.clone() + + # If d == 1, loop sets reach_dminus1 when h==0 not visited; keep as layer0. + if d == 1: + reach_dminus1 = src.expand(-1, n_samples) & pool + + # Nodes at dist <= d-1 are always included in A_t (conservative constructive core). + A = reach_dminus1.clone() + + # Force-include distance-d nodes that are adjacent to distance-(d+1) nodes, + # because those layer_{d+1} nodes must be infected at t+1 and need an I neighbor at time t. + if d >= 1: + bnd_need = layer_d & (torch.sparse.mm(self.adj, layer_d1.float()) > 0) + else: + bnd_need = torch.zeros_like(layer_d) + + # Optional boundary nodes at dist d that are NOT needed for layer_{d+1}. + bnd_opt = layer_d & (~bnd_need) + + # Sample add decision for optional boundary nodes using p_add = 1 - sigmoid(zI_unsorted[t]). + # (High qI => more likely to become new, so lower include prob.) + p_add = (1.0 - torch.sigmoid(zI_unsorted[t]).unsqueeze(dim=1)).clamp(1e-6, 1.0 - 1e-6) # (nodes,1) + if bnd_opt.any(): + rnd = torch.rand(self.n_nodes, n_samples, device=self.device) + bnd_add = bnd_opt & (rnd <= p_add.expand(-1, n_samples)) + else: + bnd_add = torch.zeros_like(bnd_opt) + + # Final A_t + A = A | bnd_need | bnd_add + + # ----------------------------------------------------------------- + # Given A_t, construct y_t: + # - blocked nodes (yL==R): force R + # - nodes not in A and not blocked: S + # - nodes in A: + # if y_{t+1}==I => force I + # if y_{t+1}==R => sample I/R with qR, but force I if needed as infection source + # ----------------------------------------------------------------- + y_t = torch.full((self.n_nodes, n_samples), SIR_STATES.S, dtype=torch.long, device=self.device) + y_t = torch.where(blocked.expand(-1, n_samples), SIR_STATES.R, y_t) + + # new nodes are those in pool but not in A (they are S at time t, become non-S at t+1) + new = pool & (~A) + + # Any node in A adjacent to any new node must be infected at time t to enable infection. + need_source = A & (torch.sparse.mm(self.adj, new.float()) > 0) + + # Force I where y_{t+1}==I + force_I_from_next = A & (y_next == SIR_STATES.I) + + # Candidate nodes with y_{t+1}==R that are not forced infection sources. + candR = A & (y_next == SIR_STATES.R) & (~need_source) & (~force_I_from_next) + + # Sample I/R for candR using qR (prob of R->I backward => I at time t). + if candR.any(): + qR_t = torch.sigmoid(zR[t]).clamp(1e-6, 1.0 - 1e-6) # (nodes,1) + rndR = torch.rand(self.n_nodes, n_samples, device=self.device) + isI = candR & (rndR <= qR_t.expand(-1, n_samples)) + # Set sampled states + y_t = torch.where(isI, SIR_STATES.I, y_t) + y_t = torch.where(candR & (~isI), SIR_STATES.R, y_t) + if compute_lik: + logq = torch_log(qR_t).expand(-1, n_samples) + log1q = torch_log(1.0 - qR_t).expand(-1, n_samples) + lik = lik + (torch.where(isI, logq, self.zero) + torch.where(candR & (~isI), log1q, self.zero)).sum(dim=0) + else: + if compute_lik: + pass + + # Force I for infection sources and nodes that are I at t+1 + y_t = torch.where(need_source | force_I_from_next, SIR_STATES.I, y_t) + + # ----------------------------------------------------------------- + # Proposal likelihood contributions from boundary add decisions. + # We only account for sampled optional boundary nodes (bnd_opt). + # Forced inclusions (reach<=d-1 and bnd_need) are deterministic (log prob 0). + # ----------------------------------------------------------------- + if compute_lik and bnd_opt.any(): + p = p_add.expand(-1, n_samples) + logp = torch_log(p) + log1p = torch_log(1.0 - p) + lik = lik + (torch.where(bnd_add, logp, self.zero) + torch.where(bnd_opt & (~bnd_add), log1p, self.zero)).sum(dim=0) + + # Commit and step left + Y[k] = y_t + y = y_t + + # Clamp left endpoint snapshot + Y[0] = yL.unsqueeze(dim=1).expand(-1, n_samples) + + if compute_lik: + return Y.detach().clone(), lik.detach().clone() + else: + return Y.detach().clone() + + @torch.no_grad() + def samp_ms(self, y, zI, zR, n_samples, obs_time, compute_lik=False): + + # y: (nodes, T+1) + obs_time = sorted(list(obs_time)) + segL = [0] + obs_time[:-1] + segR = obs_time + + # Allocate full history tensor (we store times 0..T-1; y_T is not stored here). + Y = torch.empty(self.T, self.n_nodes, n_samples, dtype=torch.long, device=self.device) + + # Hard constraints: directly fix observed snapshots (except the final y_T which is not in Y) + for t in obs_time: + if t < self.T: + Y[t] = y[:, t].unsqueeze(dim=1).expand(-1, n_samples) + + if compute_lik: + lik = self.zero # will become (samples,) after first addition + + # Helper: fetch logits conditioned on the i-th observed snapshot. + # If z has only one conditioning slice, reuse it for all segments. + def _pick(z, i): + # z: (T, nodes, K) or (T, nodes, 1) + if z.size(dim=2) == 1: + return z + return z[:, :, i: i + 1] + + # Sample segments in reverse order (right endpoint always known). + for i in range(len(segR) - 1, -1, -1): + TL, TR = segL[i], segR[i] + yR = y[:, TR] + + if TL > 0: + yL = y[:, TL] + start = TL + 1 # only fill (TL, TR) + else: + yL = None + start = TL + zI_R = _pick(zI, i) # conditioned on y_TR + zR_R = _pick(zR, i) + + if (yL is not None) and (zI.size(dim=2) > 1): + # TL corresponds to obs_time[i-1] + zI_L = _pick(zI, i - 1) + zR_L = _pick(zR, i - 1) + + # Combine evidence in logit space, then clamp. + zI_seg = self.clamp_z(zI_R + zI_L) + zR_seg = self.clamp_z(zR_R + zR_L) + else: + # first segment (TL=0) OR caller provided only one logit slice + zI_seg = zI_R + zR_seg = zR_R + + zI_sorted, uid = zI_seg.sort(dim=1, descending=True) # (T, nodes, 1) + uid = uid.squeeze(dim=2) # (T, nodes) + if compute_lik: + Y_seg, lik_seg = self._samp_seg( + yR, zI_sorted, uid, zR_seg, n_samples, TL, TR, yL=yL, compute_lik=True + ) + if yL is not None: + Y[start:TR] = Y_seg[1:] # skip the clamped y_TL + else: + Y[start:TR] = Y_seg + + lik = lik + lik_seg + else: + Y_seg = self._samp_seg(yR, zI_sorted, uid, zR_seg, n_samples, TL, TR, yL=yL, compute_lik=False) + if yL is not None: + Y[start:TR] = Y_seg[1:] + else: + Y[start:TR] = Y_seg + + if compute_lik: + return Y.detach().clone(), lik.detach().clone() + else: + return Y.detach().clone() + + +def q_loss(q_net, data, I0, bpar, n_samples): + T = data.T.item() + n_nodes = data.num_nodes + Y = diffus_gen(T = T, n_nodes = n_nodes, edge_index = data.edge_index, I0 = I0, n_samples = n_samples, pI = bpar.pI, pR = bpar.pR) # (T+1, nodes, samples) + q_liks, zI0, zR0, zI, zR = q_net.lik(Y = Y) # (samples,) + return -q_liks.mean(), zI0, zR0, zI, zR + +def q_train(data, bpar, args): + I0 = (data.y[:, 0] == 1).long().sum().item() + q_net = QNet.make(data, args) + q_net.train() + opt = optim.AdamW(q_net.parameters(), lr = args.q_lr) + pbar = trange(1, args.q_steps + 1) + for step in pbar: + opt.zero_grad() + loss, zI0, zR0, zI, zR = q_loss(q_net, data, I0, bpar, args.q_samples) + pbar.set_description(f'[step={step}] loss={loss.item():.4f}') + q_net.backward(loss, zI0, zR0, zI, zR) + opt.step() + q_net.eval() + return q_net + +@torch.no_grad() +def t_mcmc(data, bpar, q_net, args, obs_time, keepdim=True): + + I0 = (data.y[:, 0] == 1).long().sum().item() + + obs_time = sorted(list(obs_time)) + y_obs = torch.stack([data.y[:, t] for t in obs_time], dim=1) # (nodes, K_obs) + zI, zR = q_net(y_obs) # (T, nodes, K_obs) + + X, lqX = q_net.samp_ms(data.y, zI, zR, args.t_samples, obs_time=obs_time, compute_lik=True) + lpX = diffus_liks(Y=X, edge_index=data.edge_index, I0=I0, coef=args.p_coef, pI=bpar.pI, pR=bpar.pR) + + tI_avg = data_make_t(X, SIR_STATES.I, dim=0).float().mean(dim=1, keepdim=keepdim) + tR_avg = data_make_t(X, SIR_STATES.R, dim=0).float().mean(dim=1, keepdim=keepdim) + + pbar = trange(1, args.t_steps + 1) + for step in pbar: + Y, lqY = q_net.samp_ms(data.y, zI, zR, args.t_samples, obs_time=obs_time, compute_lik=True) + lpY = diffus_liks(Y=Y, edge_index=data.edge_index, I0=I0, coef=args.p_coef, pI=bpar.pI, pR=bpar.pR) + + # Hastings acceptance + a = torch.rand(args.t_samples, device=args.device) <= torch.exp(lpY + lqX - lpX - lqY) + + X = torch.where(a, Y, X) + lqX = torch.where(a, lqY, lqX) + lpX = torch.where(a, lpY, lpX) + + tI = data_make_t(X, SIR_STATES.I, dim=0).float().mean(dim=1, keepdim=keepdim) + tR = data_make_t(X, SIR_STATES.R, dim=0).float().mean(dim=1, keepdim=keepdim) + tI_avg = args.t_keep * tI_avg + (1.0 - args.t_keep) * tI + tR_avg = args.t_keep * tR_avg + (1.0 - args.t_keep) * tR + + return tI_avg, tR_avg + +def main(data): + # parse obs times + obs_time = [int(t) for t in args.obs_time.split(',') if t] + obs_time.append(data.T.item()) + obs_time = sorted(obs_time) + # estimate diffusion parameters + bpar = b_estim(data, args) + print(f'[est] pI={bpar.pI:.4f}, pR={bpar.pR:.4f}', flush = True) + # train a proposal network + q_net = q_train(data, bpar, args) + # estimate transition times + tI, tR = t_mcmc(data, bpar, q_net, args, obs_time = obs_time, keepdim = True) # (nodes, 1) + T = data.T.item() + tI = tI.round().long() + tR = tR.round().long() + # compose a history + with torch.no_grad(): + y_pred = torch.zeros_like(data.y) # (nodes, T+1) + y_pred.scatter_(dim = 1, index = torch.minimum(tI, data.T), src = torch.full_like(tI, 1)) + y_pred.scatter_(dim = 1, index = torch.minimum(tR, data.T), src = torch.full_like(tR, 2)) + y_pred = y_pred[:, : data.T.item()].cummax(dim = 1).values + return y_pred + +args = get_args() +tester = Tester(args.data_dir, args.device, main) +tester.test([args.dataset], seed = args.seed, rep = 1) +tester.save(args.output) \ No newline at end of file diff --git a/gcn_ms.py b/gcn_ms.py new file mode 100644 index 0000000..b1fa075 --- /dev/null +++ b/gcn_ms.py @@ -0,0 +1,178 @@ +from inc.diffus import * +from inc.test import * +import argparse + + +def get_args(): + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type=str, help='dataset name') + parser.add_argument('--seed', type=int, help='random seed') + parser.add_argument('--data_dir', type=str, help='dataset folder') + parser.add_argument('--output', type=str, help='output file name') + parser.add_argument('--device', type=torch.device, help='torch device') + + # Diffusion parameter estimation (same as baseline) + parser.add_argument('--b_pI0', type=float, help='initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type=float, help='initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type=int, help='optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type=float, help='learning rate in diffusion parameter estimation') + + # GCN hyperparameters (same as baseline) + parser.add_argument('--lr', type=float, help='learning rate for GCN') + parser.add_argument('--epochs', type=int, help='training epochs for GCN') + parser.add_argument('--batch_size', type=int, help='batch size when training GCN') + parser.add_argument('--units', type=int, help='hidden size of GCN') + parser.add_argument('--layers', type=int, help='number of layers in GCN') + parser.add_argument('--dropout', type=float, help='dropout rate in GCN') + + # === New: multi-snapshot controls (align with inc/ditto_ms.py). === + parser.add_argument( + '--obs_ts', type=str, default=None, + help='comma-separated observed time indices, e.g. "0,3,5"; ' + 'None means single-snapshot (use only the final snapshot)' + ) + parser.add_argument( + '--obs_k', type=int, default=None, + help='number of observed snapshots; if None, set to len(obs_ts) when obs_ts is given, ' + 'otherwise 1 (single-snapshot)' + ) + + args = parser.parse_args() + + # Normalize obs_ts / obs_k exactly like inc/ditto_ms.py does. :contentReference[oaicite:8]{index=8} + if args.obs_ts is not None and len(args.obs_ts.strip()) > 0: + obs = [int(x) for x in args.obs_ts.split(',') if x.strip() != ''] + obs = sorted(set(obs)) + args.obs_ts = obs + if args.obs_k is None: + args.obs_k = len(obs) + else: + args.obs_ts = None + if args.obs_k is None: + args.obs_k = 1 + + return args + + +def _select_obs_times(T: int, obs_ts: list | None) -> list: + """ + Choose which snapshots to treat as 'observed' inputs. + - If obs_ts is provided, we use it (clamped to [0, T]). + - Otherwise, we default to the final snapshot only (t = T). + """ + if obs_ts is None or len(obs_ts) == 0: + return [T] + times = [] + for t in obs_ts: + if t < 0: + t = 0 + if t > T: + t = T + times.append(t) + times = sorted(set(times)) + if len(times) == 0: + times = [T] + return times + + +def _build_input_from_labels(labels: torch.Tensor, obs_times: list) -> torch.Tensor: + """ + labels: (batch, nodes, T+1), dtype=long + returns x: (batch*nodes, K) where K = len(obs_times) + """ + # Stack observed snapshots as feature channels. + xs = [labels[:, :, t] for t in obs_times] # each: (batch, nodes) + x = torch.stack(xs, dim=2) # (batch, nodes, K) + x = x.reshape(-1, x.size(2)) # (batch*nodes, K) + return x + + +def _build_input_from_data_y(data, obs_times: list) -> torch.Tensor: + """ + data.y: (nodes, T+1), dtype=long + returns x: (nodes, K) where K = len(obs_times) + """ + xs = [data.y[:, t] for t in obs_times] # each: (nodes,) + x = torch.stack(xs, dim=1) # (nodes, K) + return x + + +def gcn_run(data): + # Estimate diffusion parameters as in the baseline GCN. :contentReference[oaicite:9]{index=9} + bpar = b_estim(data, args) + + # Shapes / counts + T = data.T.item() + n_nodes = data.num_nodes + n_cls = data.y.max().item() + 1 + + # Observed time indices used as inputs (multi-snapshot). :contentReference[oaicite:10]{index=10} + obs_times = _select_obs_times(T, args.obs_ts) + k_in = len(obs_times) + + # Model: input dim = #observed snapshots, output dim = T * n_cls (class per time). :contentReference[oaicite:11]{index=11} + model = gnn.GCN(k_in, args.units, args.layers, T * n_cls, args.dropout) + model = model.to(args.device) + + # ---------------------------- + # Train (self-supervised on synthetic data from bpar). + # ---------------------------- + model.train() + I0 = (data.y[:, 0] == SIR_STATES.I).long().sum().item() + opt = optim.Adam(model.parameters(), lr=args.lr) + pbar = trange(1, args.epochs + 1) + for epoch in pbar: + opt.zero_grad() + + # Simulate histories with estimated parameters; arrange as (batch, nodes, T+1). + labels = diffus_gen( + T=data.T.item(), n_nodes=data.num_nodes, edge_index=data.edge_index, + I0=I0, n_samples=args.batch_size, pI=bpar.pI, pR=bpar.pR + ).transpose(0, 2) # (batch, nodes, T + 1) + + # Build per-node, multi-snapshot inputs from the chosen observed times. + x = _build_input_from_labels(labels, obs_times) # (batch*nodes, K) + x = x.float() + + # Replicate edges across batch. (2, E*batch) + edge_index = ( + data.edge_index.unsqueeze(dim=2) + + n_nodes * torch.arange(args.batch_size, dtype=torch.long, device=x.device) + ).flatten(start_dim=1) + + # Predict all previous T states for each node (0..T-1). + logits = F.log_softmax(model(x, edge_index).view(-1, n_cls), dim=-1) + target = labels[:, :, :T].flatten() # (batch*nodes*T,) + loss = F.nll_loss(logits, target) + + pbar.set_description(f'epoch={epoch} loss={loss.item():.4f}') + + # Standard backward/update. + loss.backward() + opt.step() + + # ---------------------------- + # Inference (condition on the provided observed snapshots in data.y). + # ---------------------------- + with torch.no_grad(): + model.eval() + + # Prepare inference inputs from observed times. + x_inf = _build_input_from_data_y(data, obs_times).float() # (nodes, K) + y_logits = model(x_inf, data.edge_index).view(n_nodes, T, n_cls) # (nodes, T, n_cls) + y_pred = y_logits.argmax(dim=2).contiguous() # (nodes, T), dtype=long + + # Optional: enforce hard consistency on any observed snapshots that lie inside [0, T-1]. + # (If an observed snapshot includes the final T, it is *input* only; we do not predict y_T.) + for t in obs_times: + if 0 <= t < T: + y_pred[:, t] = data.y[:, t] + + return y_pred.clone() + + +args = get_args() +seed_all(args.seed) +tester = Tester(args.data_dir, args.device, gcn_run) +tester.test([args.dataset], rep=1) +tester.save(args.output) diff --git a/gin_ms.py b/gin_ms.py new file mode 100644 index 0000000..9ab56d4 --- /dev/null +++ b/gin_ms.py @@ -0,0 +1,156 @@ +from inc.diffus import * # diffus_gen, SIR_STATES, b_estim, etc. +from inc.test import * # Tester +import argparse +import torch +import torch.nn.functional as F + + +def get_args(): + parser = argparse.ArgumentParser() + # dataset & runtime + parser.add_argument('--dataset', type=str, help='dataset name') + parser.add_argument('--seed', type=int, help='random seed') + parser.add_argument('--data_dir', type=str, help='dataset folder') + parser.add_argument('--output', type=str, help='output file name') + parser.add_argument('--device', type=torch.device, help='torch device') + + # diffusion parameter estimation (used to synthesize training labels) + parser.add_argument('--b_pI0', type=float, help='initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type=float, help='initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type=int, help='optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type=float, help='learning rate in diffusion parameter estimation') + parser.add_argument('--b_pImax', type=float, default=1.0, help='upper bound to clamp pI during estimation (safety)') + + # GIN model & training + parser.add_argument('--lr', type=float, help='learning rate for GIN') + parser.add_argument('--epochs', type=int, help='training epochs for GIN') + parser.add_argument('--batch_size', type=int, help='batch size when training GIN') + parser.add_argument('--units', type=int, help='hidden size of GIN') + parser.add_argument('--layers', type=int, help='number of layers in GIN') + parser.add_argument('--dropout', type=float, help='dropout rate in GIN') + + # multi-snapshot settings (align with ditto_ms.py) :contentReference[oaicite:4]{index=4} + parser.add_argument( + '--obs_ts', type=str, default=None, + help='comma-separated observed time indices, e.g. "0,3,5"; ' + 'None means use obs_k snapshots ending at T (default: final-only)' + ) + parser.add_argument( + '--obs_k', type=int, default=None, + help='number of observed snapshots; if None, set to len(obs_ts) when obs_ts is given, ' + 'otherwise 1 (single-snapshot)' + ) + args = parser.parse_args() + + # normalize obs_ts / obs_k following ditto_ms.py conventions :contentReference[oaicite:5]{index=5} + if args.obs_ts is not None and len(args.obs_ts.strip()) > 0: + obs = [int(x) for x in args.obs_ts.split(',') if x.strip() != ''] + obs = sorted(set(obs)) + args.obs_ts = obs + if args.obs_k is None: + args.obs_k = len(obs) + else: + args.obs_ts = None + if args.obs_k is None: + args.obs_k = 1 + return args + + +def _resolve_obs_ts(args, T: int): + """ + Resolve the list of observed snapshot indices for a given horizon T. + Rules: + - If args.obs_ts is provided: use it directly (values in [0, T] allowed; -1 means T). + - Else: use the last args.obs_k snapshots ending at T (inclusive), + i.e., max(0, T - obs_k + 1) ... T. + """ + if args.obs_ts is not None: + ts = [(T if (t == -1) else int(t)) for t in args.obs_ts] + ts = [min(max(0, t), T) for t in ts] # clamp into [0, T] + ts = sorted(set(ts)) + return ts + # default: use last k snapshots ending at T + k = int(max(1, min(args.obs_k, T + 1))) + start = max(0, T - k + 1) + return list(range(start, T + 1)) + + +def gin_run(data): + # Estimate diffusion parameters for synthetic training labels + bpar = b_estim(data, args) + + # Problem sizes + T = int(data.T.item()) + n_nodes = int(data.num_nodes) + n_cls = int(data.y.max().item() + 1) + + # Resolve observed snapshot indices for training/inference + obs_ts = _resolve_obs_ts(args, T) # indices in [0, T], inclusive + k_in = int(len(obs_ts)) # number of observed snapshots (channels) + + # Build model: GIN maps k_in-channel node features to T * n_cls logits per node + model = gnn.GIN(k_in, args.units, args.layers, T * n_cls, args.dropout) # :contentReference[oaicite:6]{index=6} + model = model.to(args.device) + + # ------------------------- + # Train on simulated labels + # ------------------------- + model.train() + I0 = int((data.y[:, 0] == SIR_STATES.I).long().sum().item()) + opt = optim.Adam(model.parameters(), lr=args.lr) + pbar = trange(1, args.epochs + 1) + + for epoch in pbar: + opt.zero_grad() + + # Simulate training batch: (T+1, nodes, batch) -> transpose -> (batch, nodes, T+1) + labels = diffus_gen( + T=T, n_nodes=n_nodes, edge_index=data.edge_index, + I0=I0, n_samples=args.batch_size, pI=bpar.pI, pR=bpar.pR + ).transpose(0, 2) # (batch, nodes, T+1) + + # Collect multi-snapshot inputs at obs_ts -> x: (batch * nodes, k_in) + # Note: obs_ts indices are in [0, T] inclusive; labels' last dim matches that. + x = labels[:, :, obs_ts] # (batch, nodes, k_in) + x = x.reshape((-1, k_in)) # (batch*nodes, k_in) + + # Batch-edge indexing: replicate the graph 'batch' times with node-ID offsets + edge_index = ( + data.edge_index.unsqueeze(dim=2) + + n_nodes * torch.arange(args.batch_size, dtype=torch.long, device=x.device) + ).flatten(start_dim=1) # (2, batch * n_edges) + + # Forward & loss: predict states for times 0..T-1 (exclude the final observed time T) + logits = model(x.float(), edge_index).view(-1, n_cls) # (batch*nodes*T, n_cls) after view below + logits = F.log_softmax(logits, dim=-1) + + target = labels[:, :, :T].flatten() # (batch*nodes*T,) + loss = F.nll_loss(logits, target) + + loss.backward() + opt.step() + + pbar.set_description(f'epoch={epoch} loss={loss.item():.4f}') + + # ------------- + # Inference + # ------------- + with torch.no_grad(): + model.eval() + + # Build input features from the observed snapshots of the test instance + # data.y: (nodes, T+1) + obs = data.y[:, obs_ts] # (nodes, k_in) + y_pred = model(obs.float(), data.edge_index) \ + .view(n_nodes, T, n_cls) \ + .argmax(dim=2) # (nodes, T), states for times [0..T-1] + + return y_pred.clone() + + +if __name__ == '__main__': + args = get_args() + seed_all(args.seed) + tester = Tester(args.data_dir, args.device, gin_run) + tester.test([args.dataset], rep=1) + tester.save(args.output) diff --git a/inc/data.py b/inc/data.py index c09bc56..2e10d29 100644 --- a/inc/data.py +++ b/inc/data.py @@ -2,7 +2,6 @@ def data_make_t(y, x, dim = -1): return torch.where(*(y == x).max(dim), y.size(dim)) - def data_simulate(Gnx, seed, T, diffus, params): sir = (diffus == 'sir') cfg = ndmc.Configuration() @@ -190,7 +189,7 @@ def data_covid_sir(data_dir, device): data_dir = osp.join(data_dir, 'covid') f_data = osp.join(data_dir, 'covid-sir.pt') if osp.exists(f_data): - return torch.load(f_data, map_location = device) + return torch.load(f_data, map_location=device) else: COVID_KNN = 10 f_s2a = file_require(None, data_dir, 'state2abbr.pyon') diff --git a/inc/samp_ms.py b/inc/samp_ms.py deleted file mode 100644 index 4c9f9df..0000000 --- a/inc/samp_ms.py +++ /dev/null @@ -1,194 +0,0 @@ -# -*- coding: utf-8 -*- -""" -每一步先处理 R->I,再处理 I->S; -I->S: -按 zI 降序遍历,当某个时刻 t 有观测时,对相应节点的本步“反向转移”进行计算, -以确保 Y[t] 与观测一致;若与 SIR 反向可达性或容量约束冲突,则将该样本标记为 invalid; -若 compute_lik=True,则返回在 Qθ 下的 log-likelihood: - -R->I:对 (y_{t+1}==R) 的位置累加 lR(不变); -I->S:只对 **最终 msk_opt=True** 的位置累加 lI(与原版一致); -对于invalid 样本的 lik 置为 -inf - -原版:Y = samp_ms(q_net, y_T, zI, zR, n_samples=64, compute_lik=False) -现有版本:Y, lq = samp_ms( - q_net, y_T, zI, zR, n_samples=64, compute_lik=True, - obs_times=[3, 7], - obs_states=torch.stack([y3_obs, y7_obs]), # (2, n_nodes) Long - obs_masks=torch.stack([m3, m7]) # (2, n_nodes) Bool -) -""" - -import torch -try: - from .diffus import SIR_STATES -except Exception: - from types import SimpleNamespace - SIR_STATES = SimpleNamespace(S=0, I=1, R=2) -def torch_log(x: torch.Tensor) -> torch.Tensor: - eps = 1e-12 - return torch.log(torch.clamp(x, min=eps)) - -@torch.no_grad() -def samp_ms( - q_net, - y, - zI, - zR, - n_samples, - compute_lik=False, - obs_times=None,#list[int] / 1D LongTensor,观测时刻(取值{0..T};允许 T 表示 y_T) - obs_states=None,#(K, nodes) Long,obs_times 对应的观测状态 - obs_masks=None#(K, nodes) Bool,部分观测掩码;None 表示全 True -): - """ - 返回:若compute_lik=False -> Y: (T, nodes, samples) Long - 若compute_lik=True -> (Y, lik),其中 lik: (samples,) Float - """ - - device = q_net.device - T = int(q_net.T) - n_nodes = int(q_net.n_nodes) - obs_map = {} - if obs_times is not None: - if not torch.is_tensor(obs_times): - obs_times = torch.as_tensor(obs_times, dtype=torch.long, device=device) - else: - obs_times = obs_times.to(device=device, dtype=torch.long) - - assert obs_states is not None#obs_states 必须与 obs_times 同时提供 - if not torch.is_tensor(obs_states): - obs_states = torch.as_tensor(obs_states, dtype=torch.long, device=device) - else: - obs_states = obs_states.to(device=device, dtype=torch.long) - assert obs_states.dim() == 2 and obs_states.size(1) == n_nodes#obs_states 必须为 (K, nodes) - assert obs_states.size(0) == obs_times.numel()#obs_states 第一维必须等于 obs_times 的长度 - if obs_masks is None: - obs_masks = torch.ones_like(obs_states, dtype=torch.bool, device=device) - else: - if not torch.is_tensor(obs_masks): - obs_masks = torch.as_tensor(obs_masks, dtype=torch.bool, device=device) - else: - obs_masks = obs_masks.to(device=device, dtype=torch.bool) - assert obs_masks.shape == obs_states.shape#obs_masks 形状必须与 obs_states 一致 - - for k, t in enumerate(obs_times.tolist()): - y_obs_t = obs_states[k] - m_obs_t = obs_masks[k] - if t == T: - obs_map[T] = (y_obs_t, m_obs_t) - else: - assert 0 <= t < T#obs_times 的取值必须位于 [0, T] - obs_map[t] = (y_obs_t, m_obs_t) - zI, uid = zI.sort(dim=1, descending=True) - uid = uid.squeeze(dim=2) - qI = torch.sigmoid(zI) - xI = SIR_STATES.I - qI.expand(-1, -1, n_samples).bernoulli().long() - lI = torch_log(torch.where(xI != SIR_STATES.I, qI, 1. - qI)) - qR = torch.sigmoid(zR) - xR = SIR_STATES.R - qR.expand(-1, -1, n_samples).bernoulli().long() - lR = torch_log(torch.where(xR != SIR_STATES.R, qR, 1. - qR)) - y = y.to(device=device, dtype=torch.long).unsqueeze(1).expand(-1, n_samples)# (nodes, samples) - Y = torch.empty(T, n_nodes, n_samples, dtype=torch.long, device=device)# (T, nodes, samples) - - if compute_lik: - lik = q_net.zero - invalid = torch.zeros(n_samples, dtype=torch.bool, device=device) - if T in obs_map: - y_obs_T, m_obs_T = obs_map[T] - bad_T = m_obs_T.unsqueeze(1) & (y != y_obs_T.unsqueeze(1)) - if bad_T.any(): - invalid |= bad_T.any(dim=0) - for t in range(T - 1, -1, -1): - y_obs_t = None - m_obs_t = None - if t in obs_map: - y_obs_t, m_obs_t = obs_map[t] - # 反向可达性: - #y_{t+1}=S -> y_t 必为 S,否则 invalid - #y_{t+1}=I -> y_t 不可为 R,否则 invalid - #y_{t+1}=R -> y_t ∈ {R,I,S} 均可达(先 R->I,再可能 I->S) - bad_from_S = (y == SIR_STATES.S) & m_obs_t.unsqueeze(1) & (y_obs_t != SIR_STATES.S).unsqueeze(1) - bad_from_I = (y == SIR_STATES.I) & m_obs_t.unsqueeze(1) & (y_obs_t == SIR_STATES.R).unsqueeze(1) - bad_any = bad_from_S | bad_from_I - if bad_any.any(): - invalid |= bad_any.any(dim=0) - mR = (y == SIR_STATES.R) - if y_obs_t is not None: - mR_obs = mR & m_obs_t.unsqueeze(1) - qR_t = qR[t].expand(-1, n_samples) - # 观测想要 y_t=R -> 强制“不发生 R->I”,保持 R - wantR = mR_obs & (y_obs_t.eq(SIR_STATES.R).unsqueeze(1)) - if wantR.any(): - xR[t][wantR] = SIR_STATES.R - lR[t][wantR] = torch_log(1. - qR_t[wantR]) - # 观测想要 y_t ∈ {I,S} -> 强制“发生 R->I”,先变 I - wantI_or_S = mR_obs & (~y_obs_t.eq(SIR_STATES.R).unsqueeze(1)) - if wantI_or_S.any(): - xR[t][wantI_or_S] = SIR_STATES.I - lR[t][wantI_or_S] = torch_log(qR_t[wantI_or_S]) - # 若目标是 S,则稍后 I->S 阶段还会再强制一次(最终变 S) - y = torch.where(mR, xR[t], y) - if compute_lik: - lik = lik + torch.where(mR, lR[t], q_net.zero).sum(dim=0) - msk = (y == SIR_STATES.I)# (nodes, samples) - rem = torch.where(msk, q_net.rem, q_net.n_inf)# (nodes, samples) 广播 rem 初值 - for i, u in enumerate(uid[t]):# u: 节点 id(0..n_nodes-1) - if msk[u].max():# 该节点在任一样本为 I 才需要处理 - vid = q_net.neighbs[u.item()]# (deg_u,) Long,邻居索引列表 - # 原版 I->S 可行性:opt = (rem[u] > 1) & (min_{v∈N(u)} rem[v] > 1) - opt = (rem[u] > 1) & (rem[vid].min(dim=0).values > 1) # (samples,) - msk_opt = msk[u] & opt # (samples,) - # 若 t 时刻有观测,对该节点进行强制(仅当此节点被观测) - if (y_obs_t is not None) and bool(m_obs_t[u.item()]): - target = int(y_obs_t[u.item()].item()) - if target == SIR_STATES.S: - # 目标为 S -> 强制“发生 I->S” - need_I = ~msk[u] - if need_I.any(): - invalid |= need_I - # 覆盖本节点排序坐标 (t, i) 的采样与对数概率 - xI[t, i] = SIR_STATES.S - qi = qI[t, i, 0] - # qI[t, i] 形状为 (1,),扩展到 (samples,) - lI[t, i].fill_(float(torch_log(qi))) - # 强制发生时仍需 opt 通过;否则该样本 invalid - bad_no_opt = msk[u] & (~opt) - if bad_no_opt.any(): - invalid |= bad_no_opt - elif target == SIR_STATES.I: - # 目标为 I -> 强制“不发生 I->S” - need_I = ~msk[u] - if need_I.any(): - invalid |= need_I - xI[t, i] = SIR_STATES.I - qi = qI[t, i, 0] - lI[t, i].fill_(float(torch_log(1. - qi))) - # 不发生时,无需 opt(与原版一致:opt=False 时不会进入 toss) - else: - # 目标为 R:本步(I->S 阶段)无法直接达到 R,不可达 - invalid |= msk[u] - #应用 I->S 到 y(仅在 msk_opt 的样本上生效) - y[u] = torch.where(msk_opt, xI[t, i], y[u]) - trs = (y[u] != SIR_STATES.I)# 是否发生了 I->S(= 正向新感染) - - rem[u] = torch.where(msk[u], torch.where(trs, rem[u] - 1, q_net.n_inf[u].expand_as(rem[u])), rem[u]) - rem[vid] = torch.where(msk[u].unsqueeze(0), - torch.where(trs.unsqueeze(0), rem[vid] - 1, q_net.n_inf[vid],expand_as(rem[vid])), - rem[vid]) - #最终 msk[u] 仅保留“进入 toss 且 opt 通过”的样本(与原版一致) - msk[u] = msk_opt - Y[t] = y - if compute_lik: - lik = lik + torch.where(msk[uid[t]], lI[t], q_net.zero).sum(dim=0) # (samples,) - - Y = Y.detach().clone() - if compute_lik: - neg_inf = torch.full_like(lik if torch.is_tensor(lik) else torch.zeros(n_samples, device=device), float('-inf')) - if torch.is_tensor(lik): - lik = torch.where(invalid, neg_inf, lik).detach().clone() - else: - lik = torch.where(invalid, neg_inf, torch.zeros(n_samples, device=device)).detach().clone() - return Y, lik - else: - return Y diff --git a/inc/test_ms.py b/inc/test_ms.py new file mode 100644 index 0000000..8aa60af --- /dev/null +++ b/inc/test_ms.py @@ -0,0 +1,136 @@ +import functools as fnt +import numpy as np +import torch +from sklearn import metrics as skm + +from inc.data import * + +@torch.no_grad() +def _get_obs_mask(data): + T = int(data.T.item()) + n_nodes = int(data.num_nodes) + dev = data.y.device + for name in ('obs_mask', 'obs_masks'): + if hasattr(data, name): + m = getattr(data, name) + if isinstance(m, torch.Tensor): + mask = m + else: + mask = torch.as_tensor(m) + if mask.dim() == 1: # (T+1,) -> (nodes, T+1) + mask = mask.view(1, -1).expand(n_nodes, T + 1) + mask = mask.to(device=dev, dtype=torch.bool) + return mask + if hasattr(data, 'obs_ts') and getattr(data, 'obs_ts') is not None: + ts = getattr(data, 'obs_ts') + if not isinstance(ts, torch.Tensor): + ts = torch.as_tensor(ts, dtype=torch.long, device=dev) + mask = torch.zeros(n_nodes, T + 1, dtype=torch.bool, device=dev) + ts = ts.clamp_(0, T) + if ts.numel() > 0: + mask[:, ts.unique()] = True + return mask + mask = torch.zeros(n_nodes, T + 1, dtype=torch.bool, device=dev) + mask[:, T] = True + return mask + +@torch.no_grad() +def test_fix_obs(data, y_pred): + T = int(data.T.item()) + dev = data.y.device + if y_pred.size(1) == T: + y_pred = torch.cat([y_pred, data.y[:, -1:]], dim=1) + elif y_pred.size(1) != T + 1: + y_pred = y_pred[:, : T + 1] + + y_pred = y_pred.to(device=dev, dtype=data.y.dtype) + obs_mask = _get_obs_mask(data) # (nodes, T+1) + return torch.where(obs_mask, data.y, y_pred) + +@torch.no_grad() +def test_skm(skm_fn, data, y_pred, **kwargs): + y_fixed = test_fix_obs(data, y_pred) + obs_mask = _get_obs_mask(data) + unobs = ~obs_mask + y_true = torch2np(data.y[unobs]) + y_pred = torch2np(y_fixed[unobs]) + return float(skm_fn(y_true, y_pred, **kwargs)) + +@torch.no_grad() +def test_nrmse(data, y_pred): + y_fixed = test_fix_obs(data, y_pred) + tI_pred = data_make_tI(y_fixed, dim=-1) + mse = skm.mean_squared_error(torch2np(data.tI), torch2np(tI_pred)) + if hasattr(data, 'tR'): + tR_pred = data_make_t(y_fixed, SIR_STATES.R, dim=-1) + mseR = skm.mean_squared_error(torch2np(data.tR), torch2np(tR_pred)) + mse = (mse + mseR) / 2.0 + nrmse = np.sqrt(mse) / (data.T.item() + 1) + return float(nrmse) + + +TEST_METRICS = { + None: lambda data, y_pred: test_fix_obs(data, y_pred).tolist(), + 'acc': fnt.partial(test_skm, skm.accuracy_score), + 'prc': fnt.partial(test_skm, skm.precision_score, average='macro', zero_division=0), + 'rec': fnt.partial(test_skm, skm.recall_score, average='macro', zero_division=0), + 'f1': fnt.partial(test_skm, skm.f1_score, average='macro', zero_division=0), + 'nrmse': test_nrmse, +} + +class Tester: + def __init__(self, data_dir, device, model_fn): + self.data_dir = data_dir + self.device = device + self.model_fn = model_fn + self.res = dict() + + def test_once(self, dataset, seed=None): + data = data_load(dataset, self.data_dir, self.device) + if seed is not None: + seed_all(seed) + y_pred = self.model_fn(data) + if dataset not in self.res: + self.res[dataset] = dict() + res = self.res[dataset] + for metric, fn in TEST_METRICS.items(): + if metric not in res: + res[metric] = list() + self.res[dataset][metric].append(fn(data, y_pred)) + + def test_dataset(self, dataset, seed=None, rep=5, verbose=True): + for i in range(rep): + if verbose: + print(f'[{dataset} #{i}]', flush=True) + self.test_once(dataset, seed=None if seed is None else (seed ^ i)) + if verbose: + print( + f'[{dataset} #{i}]', + ', '.join([ + f'{metric}={scores[-1]:.4f}' + for metric, scores in self.res[dataset].items() + if metric is not None + ]), + flush=True + ) + + def test(self, datasets=None, seed=None, **kwargs): + if datasets is None: + datasets = DATASETS.keys() + for dataset in datasets: + self.test_dataset(dataset=dataset, seed=seed, **kwargs) + + def print(self, brief=True): + for dataset, metrics in self.res.items(): + for metric, scores in metrics.items(): + if metric is not None: + print(f'{dataset} {metric}:', end='') + if brief: + print(f' {np.mean(scores):.4f} ({np.std(scores):.4f})') + else: + for score in scores: + print(f' {score:.4f}', end='') + print('') + + def save(self, f, verbose=True): + return torch.save(self.res, f) diff --git a/input/covid/dist-covid.pt b/input/covid/dist-covid.pt new file mode 100644 index 0000000000000000000000000000000000000000..57fae130bbdd9b54f0035ae5296119047615c616 GIT binary patch literal 947883 zcmeF)KhM9}b(Z~c{yGy#h@e3fw4or0Xeba8icksyB80}Gu@c(~7FNvIlSp(lAi$+!3tOt(gJlEQ5uf6vEUcWmQQvTrgfBqLf`Q+Q*{^bAq z&)@vy*FX9G4}bdk-~Zl^|NMvF|LyPp(dU2kTmSyw{pk08_g6ppH-Gh?|M+v?`~1g0 z`QE?%t?&NX=Rf?>Pe1+cPk!*nfA+&4egC_k|KLx5`s1H`_s{>qr=NZDo4@&;U;d*0 zPru0gqd)nVJ`_G^K(ii{t?|;Ex{x9GA-j9FuqaS?l^B?~BPk;L9 zzx>JffAEtZeE*;R@Oz(s`q{Vt-XDJUE1!S%=^uXftH1RJzyHhM{;&VzfBpCW_C?Jv z|IZ)H=Rf%ImsNl7cmK}Mp7*a#{+IuMKl|Ae)Udp`Qj)Kdyjnk>yO62_SOE2fBmaC%F_q( zvA%v9=Qm#-c7J}?$y0aHyxezbUVQJRe&c<&wEe6%fA2({>-QwTd3|3Jf-_YbEvK_UYeV^nn$(whWb^H7GLzf8ee@lhx*-* zzxw8@U$zhR^hx{uKjrFEz0y13w=egB`g1sNy?A2dfA`JI2Yd&0)u zPyI`)hpP@>-eTXCJblpekni`0-ZhPUG)Lqke^FicvS<#CFJGRxIh4*{pYD44>vGlS zm;IhS)$Reiu0Fl)tG+p}du+aUhyMP;<%^>{>^<`Bzxwlc;9~zBkf#siV}1QJ&Tqav z)CY}SCr{l)^K##%dGWoQ`i=M9()P37y?8(BT)!*%&EunU+n2wnA8Z~USDY`;dHs8h z=2@*@{i}!ki|XK&f5p{#O81H8P+ecVG%s^Ck81A?^{XB&zUX~5zWQ(u^}8Q`_03nm zY#-|BllJ?6%GIZOrFX(_U+x3-=X~hh z0~h=6fINL5AM5L2FyM9mdo5x4z zwl9BCKiE7zt~g(w^P0~^{pw#m7W?;-$Hnt9ev=Z>V4EY4Jtx zt8sO@uXCs`kH*zE-@3AWsHacb@Bb;y+4_s#iM;k@E~r1}L;pUNAN@O49s6PXv%fgX z6Gy%|x##NTH;+d9FY04ae~a!3n`htpmsSr~9lpH9zAJh9pyeUodvg!C{Kn+1|9k|$k2juAk`B+~+jq{r?5A{J~*U3|N(Y*R@X!D}IkNV>N{V@O1e%8Ae z?}YF7u_(WBK03F3`HT9&=J9dG`SP60d@kx&|LP(CqB?lxUvYJw(z`%&sID*Wy;Ym5 zd6etjsiXhqReN9c_0>L{L;dc@*KgPJ)i2wJditdO{-1L7sb1-w@Y|Pjq5hl?uiwGb zzhhg+e%Su(FOKrWQQVx|bM^9@N27ff&9QO5z8BpSHqZI$Us^p}b@=iY`>y2agO-Q< z&fk3b$VWc%{r=2((K_+o6KxJOF28YeD4oAPt$sECE%l{dHNP~^((ku^b?-&IKJ7j( z>fb%m*!A8W`uhu)FOKrC_sF-u{%G`f;9~zBkf#siV}1QJ&Tqav)CY}SCr{l)^K##% zdGWoM`r_VU-z{xF>&@RgQRn(S$!{JXo!h?rMg3s&__*SHdCqHorMdC-tAF*7e^DL0 z@~^l$Pw76<9IES!m*!)x=27jvp?=k)#TUJ=##bNCp?>${ufF-}m+eD6ebRpaPr3S3 zuk=p%?aO_j{+tiJ`-}dL?L3xNAMKC&6352n@$F|`)yr=ljrL#Ehd7P;TkQQbPu=>L zZrxRfFW>tq{X16QJLAhkzTe-XdATR#qxCB;59K@Onw#@CzI=J&=1@9+eY)%Ei}qi% zUYuU^-qC!Q_FZTmAFZ!;9~brSK56WFb4ByV<%^>{G&lS5^@Z)n?_km2i6~DW$jAEn zX`J7DdD#8=T_;c7Me}n1<$5=Kzgx84{7d`w9^|Rp`>F1Fb=5)VZru9PIq;)#>(p^B z^SP*B{i}!ki|U|!@8&JNGgNoc`BDGX^~KF?Y3oY&kNPurT71#F6F0y1Wq-c<-H$Kd zxoA|U)Ca1kPulPQDfQuAFM2of%mwwMF715i-YLG^@W{Z^TbiydGues z{N~YUpGEUSd9-sax+iR&^VPq!dbsNF<*m9aZoPH#kl*>6@A4a$*WaBvm-4ak>fW39 z@*6jY()sJt>R0nGT8Hwh^|9Ew(tgJ@)^}gxSf6$;7tPr^8oS>25Br^6{>qcz_sqBd z>hJpRDSqv%4=2iRKYSFY^+SvEY4NiAi+7zobr;RcTo%0(G*24+`?20H-#+SiCw#w) z#r6A92c6qF`PP@>_|dp^>NuzQT-2}r)kFS8b-pROzjD?0F3|l}*B9^ii#AvDs9yc4 ztKa5Tdw+a-_2Hh??|%H%H(&iyAE=(bXutob)Tg=^?IX|rXb$d$?|j(bf%9IzzhhUv z{ZOB1e{pQQdETk>RWHAJG}>p;98n&v??v~8&FlFaUs^p}b@=kGdW~C;jnkdK`7Xb4 zd43OQ&ZT@5r^^*zzV+hHzvx|{yrubQALOfp^=bF8Xr86tas6udb@}ql?KZc{h9- z{kyZ?FaO%d?_|;MMVxk@=-kf9mtTtGN8{G1V_xQSQNQ|E5BV3>`KIXp%2nUHc=~;6 z9$lKNc~r0dygU6iuiE?K+fP5;*FCG>{rKv1JzstAgw_YDr%&4N|0&&*Iv4FPPv7X= z&bRWMuXJ8MUH0!-aedgIMt<42`;_NgPig(iH+Ol9);G>K-$nO@&2#ShmsSr~9lpFp z@3(P%@Z}-jJkfigk&otxeB>{x>wdhK(j1=p@~kt5()sJtT~A-M|DyHc^rGMMTe=5x zUTmKGSahHIudc6-cZdCVc=;<&e(#a5FVqM6`)_gmz1V)bUYxHVTAWXdm-?vRb@J3* zG_Um@_}&GYqd1DM_v`m$z4?15*Lrp3)5u5Xc20ixt^al7t>?VXhn*X(U;V3x{EO<8 z*SZyN9rszEcJ8Ni{)@dI_4)cj{g^xdqIcCib?oO}^l|Mczw7zxcMk2tKKi8n{-4r5 z)?f5K+rQd!1>fVcZecJq<(mgs)*NJO)rIBh+8$L7Y@hy7{1zVrR>G34t5`B+~+ zjq{r?5A{J~*U3|N(Y)MuXN_v34^&TIwBP?z>Qmi|_K|0QGzatHJ0D)Z195Ypk&ph4?S3@czieFp^32Kk zs+Zq98tt>FF3O|b!=n4b<~eu$ORI;g4!?PPv>%P;Ps>BTdHTKJ@*9`uJ)t?5^0D#i z-kbRH8#jm2`Rmi_SMx7ghw`iSv1p!Uf6wlvdFCvR_35WHXZ2`oo&G!N?=KXm7v-Vf zhj-1l|LV`*fs6fjKwkUeqd2V}TAWXdm-?vRb@J3*G%xd6G&gjweP7<;dcS;i?57@I z-lE@)xZjgH=-kf9m#+@;k&k@jV}07bomaK_=~w^iqd3Z2n(ue=doH2OQX^IO{a#cAa0L%-F_m$zsi^iCI>XJ6+>KJwcKzjTkKzWA5+`;@1jMg56a zw@-2X*bkS#;_9KkOLbA)y2YMrsLR{hpRr{-v`K7AN7Ov-BaUy`RXrPzi6Ghi)$X{g64#zM++ z*{5tB>-aSCFFHr5pS50p^4wD?k5B6pyT3Whx88neUGwCtzs}L$y*$3}6#5>}<+Y#s zi`JJtPj&N`wyyivx88T>?ccH1E!yv+T=(r=l-@}xkMG4Gn-{oKNwI5%9 zOFIY3*N-{VSbu3gI!~!?ef=!Lclm3*eAEy9 zbY1h7_MJuFR~q??^1NfTFWqxEAFlfRe((9}qkgcy_tiLGzWR&SFIuPWqB)tvqWPhB zZ9VGW`}OYtjq>$HU;Eovo_>&z`aya6!1if9c3-}FD6g8YkBhCZPZUS{RrB#}AMMk9 zp?=gaul4G@rF|Fet1hiS>uB^2=~6v;>Y(|GTi-an?#p`nmCl3Wz6a}jRiCz>dT(j} z()rb|e!7pk&WH8Y^Lxhr9c%vLr8tV)@1rz_uD4I=J(cqK>N!91k*^Qrm)>V-U;d@t zyFC3Y>Px)3eTwVHez^P*E%sb1@2Rg|>3pajt#7(<=keaKeziQm8(RHR z-ldw4 zchNk|XHg!STkGn_yIoYrdV1~8w=dmw;%J^|o!?V`ceFfIk4F2WylTGrT-1kk^kVbq z(z(s!E!`KYXC7$Z*0+vNBmd%>gMQX}_acujor6#76WfP#%eUTsXkGKQ+LsvE_S`S=@;$m--Si>tfTdz{zduNeDzVD-yizD(ehAT8ug9xs`=(~Q6JXP zi_N1;=UngR^34&|>ppbrS;wc5e{s!$-+Hur@h;tCseeAL4%!dpS!X|ULix>;ul_ZM z75AM&?~N|M?ta#nJ3(*y*u!&-%C22kNKisg`$9f9hj>b-aU%{X5pZlAW*?q<9SIaA%3&rWh<}KatqrbP7FJ8)D@0qV({g^klZnf_`imQ+M z!TR1;<9zw*FIvB7ow|$OgZV6)8=9N-sDJNuQ61}O`_pCPU9T^B`a$(pUwnC}9*y=# zdDVPgMC{L+5tT>ahG_nkuD1G==YI&Y~zR1fRZPpMA(s;{2=^>-dF-`}y~*md^%DA#@4 zr_@I&kMG8!$^N4%jbmRQBA76hsQMwoeRb3#pW&D@2~Ic^2JN}-Z45StzZ3g&aGSB z@4dMCs2{BFeKpRPul}O-i`J>TXiny{Xl`hJ)}#Kt+eLM(qp$t#E6+M?Tt6tU-~aa6 z`gExt+P|8wkBjzMb;MC#H6P#h(LVc*)N_w$zvhecY2;t*ebMe!-|Eo%(l?Fzppma0 ztxhR!edGFF_hr5PO6Ngw--GqNs!!Wby|=VKs^|PmKix-N=fnEyc?XN;UtWL5+V7(@ zhpyL0-^o+G>hjUP?gjbf+W+deaqpbpeX84UwR=W!=Xpx&^k*F(#h=o;Mdw}dx4wF% z^Psx4^U#g={@ov+R)=2HpMCJDzw*>U=fUNRWAm2w?$Gy$M*gBa=R*Bd>&Lut)#v+r z5c%q(ez3m%8|TYcf6@9y>(pH|r;A;0Zu&+0u75Z9>RCtYL;Z{LvH9wwJikBmd!yx{ zx-{w=1HbiX_u?J8$5Q`%S{<|> z%CpXX=7jQ_Ctv+*4lC|Eh29%oe%<}7FMFP)`HSXs)fYEs_e|sR#r+-IJbmc1dinAe zy@%3!D)q^???v-zA8TFvskT4rOCG;^y<5IMm-e1fe$Qop6u-1Qlp;~xX!JP_@cbh{`HsE&!YX+?Ya16^Oo+My~E{;m-6j{&P}&Zacte{ ze(%NANBv-ZeKgLOul}O-i`J>T==|<`vG2rps&)0RjyxLW>7TZbJmj|?Umwnc>Y_aV zE-iN7YWtS@p$_ujQeE%kE!`J7SNkw$dAQDf#kl-g$f9rcypH}zndysG6()kzdd(nDx!20U+UT8G`OZz)k zp8Y;b_uTdBm)=dOf4*~H?A-LR)~UbfJg9&9{OZ1wYJKvny*m`QAKDk|FU?2moe%5N z`gzNB4_6&?DqBw-e(%3==SSx&uef~s$;YSu%2NlsU(c&f*}SDYXYX+N;-!3ZM}1ZE zofBKP+IJqs)kpncee-CXFJJvd>ldw4chNlD_o6wW->>zke|6;1s2};Xc-ei$`PS3z zr~CN+FSb5i+7F$#ny-(G);ovw$bU<9y^pu-doz#r;XZm!ekrg1((B&%tw%e*zG-w0 z8u`voyN^=b`tGOC`tq$`=XV~Ir$65*bkB5YKY4FyU({F6Q!TG-9dSPP-15wK(S3<8 z&G&b#`X8l!yS{zc_bJsg7vv*fAILA=SJ`)aY41y(eHZ16S33`ix6eykr!V<@6n{$V z7JIIh_taOfbRJZXb{@L%zQ4XxS{-^(fA+zrzWbKHsO~x!zr5Co_xtF3x_t3cet&n) z#n-3(%^O>{+W$K-6jvYhgZ0h7alU-@7p-5kPTfUwGM7bpXr46cf4yJ72m4q@t4o)S zThDi&s2^N;`m*lYSKR)nel=eo7xiHsjr@!HlxO{V50~#;Xm0i?ThBT^jr@zwvAE`I zAAQrMbMUMA-QS$$qy5l2arMg1!MXaoukSmBz6W$^-|qKRTfeA|IBgCW&CPr-+6V22 zYwnH9$9|P*}mw-)~zZ^I`%J;wTe9`+9ukQc;bn)(6ZGZhB zAFXSgkIkq1`=tF2X0n4e#Fhe zxoDJ!@{o_q7e{&MT_L};9`(8CK5^yMZ~dkD*gCZROY4^3^H+N(Z)wix9kxEL&uZuM zPA}R=o^z|Oo_+Mkx2{yz9BAY(Hg9SEqQ0G{eey4=gYx_Rf2w^aQ9bobb4T&&`r^)A zZEoTht#@y0KY6Gw?fcJfAFVImIa%+!YaNPT)R%o|zlW!^&b?mrUgVhzI%b>;c@^PB!VR{hc(`E>i&=c4?@HOH0D*B9Df9QpFRpGEIVy!zL^ zx&Rcn-=h8XgM74p#pR)V_p^R?ejkl5U!J&gmYtKhby&ak<)MEs(fP_NF5f#t^DV{u zyZ1f0Xq~=F{h&B)F36YX-ccNvFOKriJ3>BMSL$ogIdJ9Gw-0}5KFTlcUs|_(=U;T+ zPw9QNo_DApak_LqeOrIg{_^^JQ_ue11K+wubFdHc7n@hT_N%XN^J<)bQ5}@;yZ@HH zlc=8h7n@IG*Sj}9ZLa3RM?Ui1TkDCJUGMuZuYI(>c+aP<@33{KerX+w+lThfpVBqYNIp1Gjkwe#uE`QEa1)LS$MbUy6-nkSCd`TG^WslQ|GU)r~9U-V+@R-0qjH?Mm2 z)jW0Odnb#&=i=4g9g26~YWwR4`Dk6^d~81L{g-}^PknjTIcM29iCc&DTVEdfZk5hi z|I*$Wn(xxS6Y~7-(05_+>POrhoQprtPJ?h{vD{nlTakF7)7 zzqD@oJ%6=#@|Na|-eK$0`mA<7@ARU5e)wseCtYe&4EV#V)K^fFY4QQ+9&^_ zIw;?F|1EtdQ9bobb4T&&`r^)AZEoTht#@y0KY6Gw?fcJfAFVImIa%+!YaNPT)R%o| zzlW!^&b?mrZseH@I%b>;bYvwZt^thuP8erb+;y8Y{O zQU2nZsMSJ%6C8Ocjx!f`10k6J7?KB ziCc&DTVEdf_Y$42yyEh`Gc?~)yuW+D7r&oH>-1IX2gPY~LB2fqj^em{ag>MN5%STx zQeTVCfh(`RefUfBQGRLv(z@k4|DyYTO7E-nyhHtn)1~w2+xm<4m)GB$diM7o_|`3& zgME;{*u3hsUwwU>SL6JP>Y#kz{kQa;MD^6a*nAqh-o5c@b2S$}@{#Y}T2H*}df$I} z?W6U@dp>o2f2~9HOY2bFKD2lKl;+`HFM2of%mw|folk$xhy083og1yA(Rs{+#`;UQ zo_{|VpZ*Tk$AQDhvN2Iw7-6kkJhia zJe2Q#*6+^0Q;jcQp15v@Oz5vNP% z)3@~(?JuvtH}&lAJ@BntGza@2f3bPhYrp#XHm}C{7u7-ezWZkFy|tcr+4a8v^4dr1i}!r$`uX+7`xP55v{3*@DywC{vG zzZdjfSiJfXHwWjUQ69=eJ}zGz<)L?l{L*^V=c4<>l~=#@m*!*Z(DpB_TYk@9?VY@( zIiq*j`m{c)ozFYHXdijbt-gBp(I4NsQeAVPk-ylyrTL5ccAoahzo-t%_uYR>-$_(Y z{nFe~yt=-)b61<2_(kj8+uBbas!RL+^V>)3i+4`e`|etY;urO0AKLHXDXnv_7rh&K z=7P@aeEM@f{}%ao3}0UVm@uIj=lEz1TU}kKf-9U!FKEzV@r%xi#K7 zF0BsA^WA?--$_(Y{fo_~vFqI*pEg%>;iEXpqg&6q(mKBHKdqiT_pH8n&nMqIDb+94 zLv`#+`&~Vye%q=GpZtpRYf(UL5)Iypu)WbMb2L4#n-Ws2}G=K6X8SQNHte|D}I_p8E2vbI!7J z61NWPx4t~|-71~4{-w=((H#3d;9tKx-w9m(=pW_L=7Rj@)o+|$v>yAu>eHn@7u_eW zy!!g)FU?2!rTt6mmf!PMyYHv;zB+I3vh`_wRy(Kl-s@uXoLhbM?4wWqRoD6KgZ#zj zEzMulxAU}L{zY|AzVH59`c9&H>X+t@;??!Vox9rH#4lR!-qwEdP+i*hpWi-OpYEKj z_j{9v;urO$58ChHDXnv_7rh&K=7P@aeEM@fs)O=<_usPLNA=V%#nJs#*B4i>+FZ@0^W>wrb?WjL zyWaV&S5KaIA-J<%<=gad>7X2M1UhUnXc>Ax` zhw~yI*E;c1zH@v3rGI~(`tq!E&a!h7w+`#KzC3harE}K5w0SR@W4{M{zn4Y7FI@fT zALY^Jg8b&yZ=7DV9=$6XmoKi*MfZs-ufD$dOY>2FY5&r?bQdi55)XL09NUp@QC<6B?azm&&!e&l2QYI%!2hyBqxE~%@Kk?bkk9zuLamipO@2I|Ys~`QNJlb54-@N*b(~H)l_k{K7QlE?N6IWh+ee;**qx{nTrFF~i zd8^&`Q+i+KjNV!6)B3D-PV0RS7MthX>Z@lT{qe85&SxLwFE(#!{-VB}r~UFTs)O=< z_utZY64g_`G>*cB6{b*cy{L zPh9KFw{)L;x)ew8p1)c?zgpiY?%tkKKknH&K8io3e%{i4btjm@lw9uC3@Gie(m4+wr;iWG>WT_`oa3X6OHrb ztG{UdqIK#nnv=OK_MMn#ef2tD+B*AKPmA->yOp>6?qi;4ee>+gZ+*Jd7ph;)H=m30 z?MEa3qCVtVf9=2G>Y=&O7xf{JM*cK5?xx-_pBY{uS@}tJURK>l?-0*Sg30bl2%0-A`#9itA^w=URDBef3KHqk6RS z(~UdVqV?kRqJ7HlD_-CE_>0yrw!f99zI@tu7JXl7Yjy?fd|)~SciOE0S9{G~ee zY5S|weU{(zI2YgdV12LZ)ArN%Tekmd_4w6%TyZ|OZ|lw1?-`fx-H4-gsJ{KHmoIP8 z{^;*m`|_*#?*F3qqEB4w%y-dw#Jz92alU;QtrM?qzvAw#G`IS6TwYu_Y-&yqc7LEKxd48AZ9%=jQ zk9H5MKHqm9`Rb#7u)g12<9zw*FIvB7ow|$WWG;*D4b6{6^~4wDm-6d-xB3(3)7Hx~ zPn750+?|@?OV+^pNp+;J&L3Ks`+UBdJmW1zSXO5zvYX6lYJ8lz8C%8 zQNH`4#qDE0)ytQ+=-s0I{2j|jaoRh$=zZuD*E;h>^Ie*c^3hzy7n@hD9*VoKr?k#I zk&pa~?!&r8=U(x*zItWnpf2Be=*FFA(Ry)uQGfQqr@r5hdtX$~xk~-juhyS@+V>QF zPif>Y%JaKKzf0Qw`lHQ#)#v+v?}dEzQ9oGUccO8=eDxQtU$jo$MRRiBPw5^@@5`JQ zyN;IMcd4%a#QD~_S2RzQhxSE&>(izEQGeC-#TVts|r@H%_SI@EPb$zvU ze7bZFbl>Q{y?a_+>)N;9Lw(x%i}LMbPBc1~d{keY?s=Sx?|a}og}w)LslNT*()!YQ zI#;^+wEb!1W9KN(yf1n$C|_M#+8~_rK3$5V zIK9}sYI8(!bLqYD%azARbAC$wEIRj!zxCBCy<=3Dc0RiC&a2vbKD}t4#-I8ivUH2f4^=a!Cd#;uD)K{-`4pfhJ4!UvYk%#qZbr#hx z^?mt%Kkj`|z14?v;L6jFJlb~_{k=sae^H*_CHh^`_SdI*dUX@ zm(8O~=dJ%#cVBb7&S9PXtLuxG`jn6E8{M~ePpfBL`}TXNPj@|C%D2y=^TL zx%j>ZzEkLXK$rTn-&^Xhbe_(Ye!7pk>S2BLyn{vG{nC53PMo%%eX5r)Z_z%bearqG zYaSQZ+}AqqVzK8jKl$P`uKoDtxU_SieEXN?SD$WNpD2#*sq4k{^OojSUUkf=Ty^-) zN85*XzS24CU)p}9{Vspyb)7!=PuDeXY2RDyceuWIDbG7b_dq)bIv=)fwSN~-Tz%9J z*7v>|=gU`r(fUQ})Lk?ub67Mt^nTqh+Go|_%eOx*ZXaB6K6br%qCDRp=Rkh&59sbsK-`_ni7&>^}6{e&y@4-%tPj!1BdQ`F@w^_etBo{dnsy7@G&b@lbZx1Rr&>Y%x;I`Yu|Uzhep_m1w{yQk%2`}TXNPdkV8>m2gE zC;6p%@=;#v_dL$U_dQtOtNOJ4?EjX|f%;abbEThB9p^;rX_Q~;uRbkaiaXa+nydV3 zzH^k$Q|gEB?^yqyeN)sgUhi36X`S=>_i52P7O(bhQQYqV%?Fp?_{v-N!&gT?i#^xM zd+MuKItQvp>z8iaxk~H$4(Y{JM?I9M4!^(4rL9}%T7J(@mwn$$d+*}aei!K9ud@5= zAJtu4b825*-}use=UDqKzx&El@1pxg>yeNA&UL-Bwa)KN9qZe#`|O3Ec!jI zIg48_uX_3N^t0G=t-PncdZphrsz*C7-MI5Cnu9v@qIZjL^Sqy@T<79n>^x}i`J&&0 zylUS^^!xL>MDx`zjp{CTUh?!+it9^!X+Ey~)O||z%%K#&_7nH+>~DVPoy$Y-ly(o^ z1IjZm)Te#xtEUg_dF46JqW!S@@yq7XXg(M9<-AY1*4ba&cmFN>eRQAdm*)MH?(Jgd zp$}Z&0q4iw-=&@3{k4ytUq5JHly6_}RsGIEd{G_spVB^W--CQ}LBD7Hto!hLzi58u zg4WTXoe^3hzqlSS`GyxO}% zar0R8&exp9t(RB5e0lm=?73FnQ(wKZ-(B_i&PzA$Jd5U_4!!8z;@dp$=PB2@_!m15 z+Izm}_aLv@_YwQMq|H~qG^)F}=GDHszVW5`&aw7ee)pB9-bMMi>UnqK=7-+7JoHX! z_uxICJo7;P+qb@Y`at)O@|<(ge%Sr^W%Fn>pNsl(-ltsa>@V)S|Caqex=-~>^L|SA zcCqu&2ln4#_wW4J`@6LByTA6)^Xmugi}LNOPxU(o@zQzKu^-xxF75O7J;*l~^n2Ei z{@SEE$kN9z~eC;gPxIj{Gz z=zAev?VX{xcZ}wP%Wr(;t^47tqo2i|Yvn!l)hnF?)uZ)GH||`eb$y5Q;;N$_%2S8m z-{sQQt#d8E=ch~W`J&&4yy|`r{qE{_fBmDni+*?ImB0G7PQ2_smsYQN>UJKqxckPw z3qFc>uIrtxb$)N^Sg$YCuXEOSkJh7g-i5sO<2-!(T+|1ee>GoU7v+ClwtwgJ-G58> zh1a~?D_U1wU)+6FI|u3m-FNR_++4-cydXG>q4mDQE6@GNNA*knT(sX)n(y29 zAm3a{=h1KV^*azZ$Ie5%zI!gsqdqNOUh(GB-ow&-=PRAB)DPc#Uwry^tb0Q56Wt$u zQ9hbS-$}Li$glS9P`vNJd{CS&#j){qKYaQ6S?sx1-cw(_(m7B)TEBGT&Q)61cStX; zI_lxdTXmPe&b9oWpDwR=FYdiB?RSCe_uKvTkM>#gok4l_)nD28(tO_y^`El)TBqK{ z-oJCAxOc^0a}`JHY3uynXz$lNQ2)+Z-#uE7)}g%iqc6UF-qPHyqtX5s<*z#Yvg^dn zOWb$=E!`KIoBE}^r}Pdk+85m$y6;ll92&p0bGSeItB2;$--~$Jds4r1P`A`KUtefH zy0p*R_aNU~O6O{S`t&=!XpZKB*3l@xbpQ2f@lw2VyL9JXZM}1p&Qt1#zy95;@9$Xm zQ8r({=sxL-^3i#{n?>(PyxM>7kK*RB=$)@Qi(4*SVJ8^V6m8*G2DMUiJF@@%Xsz0;mYpUx%D@BaF5|0qu%i|S+d<(JK) zQC}C$*E$;c<%;W5+;{&i-4~jl`lWk)O80ZoInet;_wC)&`nAscc}nMYf9jz7FbDl~ z-<7Wq^*aaWxu`zsi!SZ+_C3fqm(m>c>wIWF7tPQ20j;C4=cBJVh?myCr8?GC^PR7B zzS26r_ifI|N4~#f>x-j%!sUxETIW4@Z;O5Z)%~7`yMHtvTz=y#Z`}`H9sMji_lm#u z)hnF?)ur`IH||`eb$y5Q;;N$_uDtr)N1VUTwfvr+F1_Q6e(&F=`p>mSu! zG>=kW*!A`GD^Hw{-KX06MfvJ2nwNZRUB34rzUC^9*3;Jcz0uyUzR-O-M}7BbJvLun z`_UKQy3&1pl=i>)(fVE2{++Ym|EJo0;WaOLPw5?8bPm)9dLPbDi<^^oh4QezbGSeI zt4CYk-;2D9>UR##bJ4!&+;nN5r|+P7=2AMReyjaXFY4D^&^j8Ow{$=CY4KA0E!DBE zn(ut2^Oe@|`%b+-Y@GIYth~~_p*a1N*17+_lWOmgU)}d7?*7qyaQTg|ymdc(b@a2? zbFI9mzIvr|pn9}^>BgO_w65=vUR-t5!MSaMt_Ip6T6TeTo z`|BUoT{IVzXJ6DO*7tjFoR8`+TEA$WdW+^GA6u92{Nih_;%Gf>o!=Yn{pt(dr*qVI zkJh7nl-GXr#g|{Y&yUjn7uB~v{g&3FbN2iHRJ$)UFZY4=llPR~!A0jleW3kIadY&p zO6TCy&f)&-uO2PGzZZEI)$bg79(DM%ebKq;(mqe$LG#R|)KB}P+sD%S_q#^jqquo2 zdgp7-;?~QnUcNm2EcRS0@2Rg|+25mjeCMSbcb-LaP={XhZt=?VUgSM}Z`Q4I@h^5B zwBIp$-?)5nl!tzwG@7q|X}soQ9(d)8>+{m`P=0Cu(z@lFL%G&n`-+<%dS~7z)^`uy z1Ip7E>f65c)zb%B=UvF_oT{zE?#EXbZ9d;j`mwPOogHJoZzU+_kP(Qf#*B`21u5?6{-?#5UzPX^^ zvwqqi?fqOdzs|$DYXflKjPKi9g3UBqIbULEN;EL>gCJR&tlKD@}BzYmAX*j#yW@9We0|Vm<4g1P-8yvppsjD7dW+Vn zgX$pPyW^W5dgt=cJEh%&_ki;Bh5EN|ef9K#-UG^Wo<;j%_v4q%qtSdW>dSdc@rzw& z{=WNf+3(}JFZ-JJQ@Xc{>Z5bv`ggqZ>iyB~tJ?Y9pZ(D~)DL!ldG_&6)$bg7F8O@A z`&`;SZ{LG_b3wmn=W2iY^nNaypShrQH1e&hzWmP1diVU4?pJ;_U;WZ~O8xWAchS7X zY2>57WA%rtfAK}@%+)(t^nS#vy*m^)k45i%%~{-fdDY98r=P{1Yvn!l)hqkmRgdqy zbmPvmXb$Soi{351&GUYqa-EBRvGbsP*U{^n0u}U;Wa!es}!Ni?0v5YEL+@N3dZ)B|@E%Z}zEJ=6t*@Rw(0f36&a-Gg?0)>R zc{G~OMSVGMDSol*%-?taE&F|3_hn!6eoFUtQGIkS^zU%ldG-Ei_f_ru?$7>c9qI?W zzdZYRr|NeOJ&$}o-F+@?pSSNpzPX^^vvajSeR@9^&CguWIvUNj)K7g{ycECYc*VUF z=aJ7}l!xty?s`6o_k8@)-?8rRo1%MK?7dcduPAP=XkJ)z8iadA%R3Uo8*4H~W;`xAAIyFY3p-r}dk+boaH-^2JN}`bG1it<#TnD86VN z%2OZJ!}{*CalU-@7p-5kPTj@M$^K>cQO`Wk??>Im=F|4^@7q&4ul3k-$kTuOTfRK; z#pYF;&&8fwo%$cuzWZ<4zE@xT*6;d{YW1=A;okcWXmhi^Y@F7w_u6&(;CpY@m-0*X z?T7lmXg%tKUhI86)%q*l!{X|rzW!TZUj5FKb}w`(Udp@Xe#OnfdF1mKz8iadA%R3 zUo8*4H~W;`xAAIyFY3p-r}dk+boaH-^2JN}`bG1it<#TnD86VN%2OZJ!}{*CalU-@ z7p-5kPTj@M$^K>cQO`Wk??>Im=F|4+zw?)NUhA>vkf;Clw|sfxi_NPxpNl=OI`u!Q z`~82a+xO~=-@0A@QLR4qKHPiX0c~#9myOf<^>*}|Bkh9 z@2k4^THSXkZm!t5)nA&A>a6?W%hS)I^RD)W4Stxq&BzZ2TNsLoSbht}B_^@;WM(>Pzg`is^tTBq)!@0xuV zyN`P2fqp;gqI|mfw0(Rp-D5Q$o!5GF9{bCq^@GcATwd9{YV*11T&NEHQL5AL|5M$* zoxlC8eY*amT7A?<+57h{+V4}XU+>laXn)#!qq~p%QhocO^IWtZ^+7K_eFx3cU-`Dr z)tCNTUw(b_F3sP(#c6SV>HU0^=HWc**dOIBHm}+`KJt4me(AegzWqDaT-9-Ji{>L< z-TM|d&+6qjk1jhGT3s5uzPjgX{H<26)E}x#>z8iayD6>nj;i^k`lY_>FRkB2{fMhW zWBsN1XgyxOcxk6U9*h_r^GkF5DBYWRsbhbXx7a*=@{x~x=R|(#y)O3eSns;|e7x@6eB@R4zQxV6 zdil+x%XJUxHs1BsJy+vzwR)xgP+eNzbmM&wv~}K9HNRB9)OY=*^}DDaadl{{zce4! zNA+mrFUm7l?C;XPxauw6I+SNW)DPBopN;e7tG{UdqIK#nn$xPoR}b~+-J$P_x@dpe z9Id0(_nnhp%}49>i~2-)w7zipjms;WS8YBQ`~6qPdCddm(`Dn={wwajaP?#V*0GLH zBmbg#VDEwMJ8=K{(|_4GtxvQ++E>1Jtv;WYU#j0ZIL}4=aO@(U%j&Zs>^rZbmQJlXhRS=eReMX9r3+VANlH7 zPouo9qf342Yq5Eqi*@yB@x|se?k$2BMO zz;z$J@AbZ%1D&I~ecB)OVZD6XJ*X?c?D=WuxoBT>Zn{+eqij9;>Qi4A)lnC%qtShs zbM@uBFSNe*^j53mJl69|@lyYM>uBU7-+JVi-s__Ipuc0~^Q*6S<{sr$_q!l&uGsU` zUz(5Vm?PGw^|RP>t-PncdZph7sz*B?-MI5Cnu9v@qP)`n^_O1#@a?Z|&&4mBw{-hb zXZhl#e0^homusE%jaNGdirWv>!}{jhIA6Z{i`Fk%r|zOTtvFv_Xq|n~e8f>5TAtq< zt-kLa_B&YL{+)+t?u{#sdiu3d({VD9iR!qwMe`A_?tP1!XZ7-% zN0*%otuBpSU)^&x{#L73>JQbW^-DMI-IUgON7ejN{Zilcm)7s1e#F(GvHsHhx9qv9 z^?~{}H@f@EM|Btbdy!{7u6_9J$2vZ({-Sz|)~UPbU9LD^Uud0u(42g)(0$PQwvJX` z-bHn-(=XPi^=~~czj1kG^Qz6~V$ZEk{Y$TP*ZwQ+zEJ(e*0GLHBmZLe>pfie@88kZ zxwJm9`TF2{*Y3%_wER;2&cXY=*mbL(yelqm(fZOoEH-~>bL{@|>vx{Cd(jV{7U!4O z+^@KOoJT$%`HRhKpVd7Vjp~)&>r;By=#+cYg8yUg+-UTzvN+pH^RAPidV# z(Y`28|9s@@vy}HyI`74vyL~jiw0}2HoqqqH>h^68>Z5*I*LpsU{EPNQ_jb{{(BERu zLF*IkZ?5)N*LxG^)7F>jcMkeT`<3=Z`E;6kyU$(F6_OZ17eJ4;pjrK9u>dSXO zX#LlveVym3C$DTBee#iyeD4tXrFUOiSIXnl{*L9Zcf|LOmhS&P6#3o*c5d~T=A$~j zullrp7JIIh_taOf^m{<{Xy>CFcb=!`ZXWGjsfYG2#ZjC_^|Ags7r*TJm$uKv_Fb)S zG$+3kx_zqilxw|xvFE9;A9><@`s&O2MeEgF^j)(rHlJ@l@y>-7U-T~g-e~oG_ny+e z&W-xkhdkQ(k?;MK@;*xYUu^x4_OaIW`~OtCFI@A|zdUr$zNPis`ts{{p0v5^hfg=omuHS2 zrT5@G>evU}Z}sJ?kMfbPF7nHMf6a&gRLk>stn)9Li+HtrM{)hrxctWD*%!^Pl#k-- zE~>lYZ+-PjeWALve(A=&5APrASIa~9-~MR#NMF?FqVu3Uda-#+^HIMv@)zZ`AM>m) z-@W0g&o^J>uln{8r|n0J^J(=LtzWcG-NiL0ePH)d&pzni5$|^z$<-g!A!dF6?t_u2ld*E!`as*n0LM?NjiSJ(c?N4|H5 z{IYYV&E=`?{dhOu^xv`Og8HP*i;sMif5n^ce*8PL*gU@*`Qj+R{Q+A!Yi>n`f zp*-_g^j^?>m+rgqj=kGO>-2~Eq>-<0TAWY!`|JC(pH|C-?o7-kW(M-|rvA(cG)O=cj$>-|x@6qsU%G!K3eBI7uWZyKCQlcc22BMqdfHcb{_MuHiwIyXY<87FXt++ zJaP0s+kf>sr@Te=QNQNMr^Wf|+8_DI_wJBicFwf9Jk`A)-w*tz{*E;l)F*9TeB`72 zE8cwf z@~!Ll*WX*q7cb@acjsJu`=dUvb*p{nQCxl057syT#`*HqU$lPFI&~M#$$dYi_hz2R z_xne2Ge?Uq$Tu(K zm*!cT3;)vYO`i8t`a9Np5~q=0nhV;GMn0M&@{y0~RJ)&x*2(klA6<$&|Du22P~Osf zv`^W7s^vS^qWAKY-i>_k*t=b{PM`Xrv3_;)#QXjAeOX&R5s|$Va|;A-^=w(p>nLc3<+=_r`njckGoXU!J+3{b=N) zIU*nVs7|%}xoDj{zZ<#~cTRNv#pYFaoqfvoQ!U@Q7QL6J^ls#P$KLItb^6o~jrFUW zC+>TOzQ;837v=SL=UjaIqP}p|=lgpQ`Rb#7u)g^>&X=$LqV-S?)xddC%Ri;n-%n|tC~hC;piy1a2aTPpJU)%qA-~jz{z`EaujZrs=y&U>_WeWm z+5YXPAGCjUeQ|XzT95AA?}OH-Joiw_q=~_D4SQ-2?K= zHE(kfzjW`%_v50!W4$x9uG(BsU$lAgk&p7Pc=O$lf6uDxH%^zx5na85{g66w)-;HY9-%shinJ4o7{!tvwz1n+z+K2x8d!x-0y-WM(mqv9_ zKQwl(^7! zr93|EUFpaEC=dDeL%#an3#yOSInPDkDfB&{OZA=aqW6H!_xpAp^RITki{`cR#GS)? zTDQ_{c~3SG@V|$M2!Ke&cj$opUal2R_Y{Z=dS+Q*E7jEP5|!zDxJrc*oxDqIIj! zE3fy^e7|q!G5>0Fx@b=F%#TLCxh-CK;%Gk3VgA+2m$#@s>bt)ST3o)m z_D4SQ-2?JVb1ltWgc{S#oed#p?$?KdI$QGPnYUr zx%U9R_$Va|=Kz?birMd7g?f&F>Ka1X# zcy<4d6*m{OA8lTI@~!jt75a|S$X}G#-<@;u?T`AvRiFPGU;SW^uRiJr>zjY$ zeEI4xTEA$Wx{K!IzMs!0?a|Nh=+^F+RVoP$PnQ9m?xuJZUa z+7IowD6ee3IG--v$NG1HkG^-+eD~>mXkYP*-huw))1~^@xcm0|p!Fx;J(lwLwC}t1 z_D6Zhmydk)y%&^MI?qMlDfB&{OZ&Lji{1muHy^)e=Q00kbGX=fu6NFNp2aIq9KFZ( zU%k#LZ&7{Jx8DVy7U!#Lf8-$C=eB`72 zE8cwfozI7;0+m}YZ{he!3o_Q>KFKE7{c;Ah8?AtLillV)0^fGk($y^v>DKlA^}a5hzZBo^ zPG5BUqdKs8Yxlk9mmk%G)y=PiL3T!#h@b2a&%v-FfN{?aP1BJy4%` zb}2vB@4S5&N_S zyRR3`YsK;F!+lzGzb?Nx{Ww?WR(t8<7Uf5Ed9Ubfe!9H&M>^8Y3+bhKmgYjgZ09D< zz4=P_tMoh8T+n`Oq@y_^9qA}fZRc~*I&uBoNH5!WZF$gq#nq;xb!F?QEnZ)X?#nCP z8}aV3d%I|zI@QC*>a`ok?;XRwZ`JurasBS}MYk`i3!Arg|6cIRkLtnd=3hTuy!?yS zFIp$>qB%M5SGrHWW3--)=Dz40ychCZRKM@1G*9HWk3QHaFRFu$ovS!H8?8fnsSfp( z{K#LMj?Tk-^_AW~be^r>e(FK{*RIYl&qeFedHX)t>J;Z3N^x|y`?Km1hvMvqbot#E z$_oekq=eIAr@t18q%7b0+{Jv5?>)2?2b7dnP ztGh?;9nwqwi|$3KkELJ!uB+YkwS5=J?{}>C>7scyj*k3nT)K1pO83n=w*D47Kejw* zJzIWu{qC7Ktj?BaQ5>4%va2_*dKRrO?Th^E=A(=Ao?Y}#i>vKDMBl0JlD_Vly5wE- zdl1FRkLtwg>Z_kFUj9Yv7p;?b(VTi-wVfN?J)$H1qImhxdHDXY?^JdBHLrQ1eb73j z`;P5{bfjZ-w*0S@uYJqMkK$_6(R%h*%8&Y4dHB(M|4RALxug5;ezWt^DEVj`eVyi+rIj|Y;m~a?Z0S!+4%aM z*H_AG9UJxWmF`#Xf%`x&`7b)J#h$DF#Me*1*gUo6MSg$hzy6N3Zqb}ysUGLcMmjFN zesSfRr}bzZ+x=!Qzc_SHt*cEhU*{)IUUu`TXVLmad94@6#>TDRUDdt2^xEEQ^qu-H z>Gofgcd_riIQdbXSY3Vf)5XibX#Jvf@-CW_b6Rw+=pJ>R*7KLX4|&9+_oLrq-zk2& zeVW%i(LQJ$($#4nq$3@xv*lm3FUqGsd60h5`HE{^`YWwR_1LG>4?3Un@{1SOIr5|X zt{yhOy4$z-rR&?hmj~4)&lP8#yr@q7vgOqe(%tj*UR7tSvwgqX>PB;5%TrrlW&38A z>ZZ$Me>xlKW&QGrb6#I5uXSv+&!Rr-r#t_P?u+xlwNAZ@&J*<`o?hF&i`MbiZoT~e zeqevcnioIPk&g65an47-Uuk}&{ua%t?Eco%?Zd9$yr2baCEU^gE1=^hI&LOLUHG^HE2+^3(kuM7sQ_9<1)WtDi1j{zdB-t&?}roSe&I z=OxZuk>2xX`%dI*yn8AReOJvx7uS5|iPm?2>*(?|KOOl|{@Qf=T(sW#Sda9tly}{S zOYeT>P+gw(X}+bn>dO|7?w@$)tqwN7y3GUi!A3gT&v}%MvmdIPZ6B0hz1V*AMfVogZJ%%YAVk zxYnz8(RreN#M9YGUu;}$btAvOAEn>1=2btNzvM5)(b?+xO7knPJZp|C58b_H+lTG` ziNoq_>lV$YwEohYS3O^8KeWEK{i=Is(fi6q`l2}BCA!yab5REySAKfmd%FCn9<1)2 zsGlxg{zdB-t&?}roSfgHxuG~VI`50tH_kiYJ3`;5_@%2ue)B}@8z-LL{OnR)D1UAG zE3J1vUB{2&YSU5u?5~s`)wS~QqdAwCzwxzw$LPG>d$#knpFU6@Y^>fq@}YS3m^-SU zZ6B0h9JU{Q(Y*)jy{gW(pSr(N-Ppe5<7eyZqPaFs-TbxbW%IDFIQ4&J=fOsGFS?bU)k|T{MP>T#`%4P-cdHv7sdH5(RazVKenHh zpYA_rL%RH^9<1)2sGlxg{zdB-t&?}roYwi%-6N!Tj`fQxt#dz+?p;Ffp!-|hzVe`X zB3*pzqBlRgR2QnRHog1RUi({*s}8#T%eCIRi`Jw1-6OQGe$f1Xr9Ae<=5yZeJ==Qq zs}oy)b?e+yb}5g#>Q}dU+Xt;vFT4Hdi|#$}PNDaJUFt_YU)lcDtsnEKt*_GhrSq$s zE#7?TY^0a<%O{TiE4!YJ_F2?t{q)YabLL0>i#^xcy(i`(j=%P`uXW;T>jU}yeP1*$ zzhlMG%gZlbob%D|S9X4E{VkeP+5N4j%g?Ug{am!3pS@^aXl~1Lp z_x#zulNIMai$mX4^U%dLuX&>N>K2c5aq=P^=~$gDf4Thnaz5ga{+05s`*7*qPv6z$ zX}_0k{)=mU@2CCM#TKV-^FZ-zq@z5}1Jx&P-Dh#usR#8hjxDc#knWy)r_g)AF6Ec+ zEA5Z%%N+H`wr|)GxzyMFgi99Cz`vuK{B z`O(=(w-1{CqPRutR$uhSt>0bMy}R_<-fMKPohRM?i}Eh^y=SWnyS}=5#PQSF@-JGy zXq~)^=C$TRcaPZa0~_h<{?f&{ALyR?j!Ns1pWVFe+dhq3y1Xb3=~$g@zgN1y&PRUa zM{%|3XkP5Elpocz^1Svd&ON$luGo0>sFUp;tDAq(yy}<7{_^R=e(OG4ht{L~?Mof< zBi()Pzpr`^s1D#%f{EQ?yqd$Y;@jV>3+SQm;2$I zaIIJ0qB-z87k2&hi_KG8UgS3qG*33t(eK#hr@vA?U)lS|zVfVjT95Lu-Ea2ti$im> zt~R~2{*~5?Z=dR2w60Vy*3Y)TIJVzu=yw5Jlgm*~4>+aG%mSAM$pAL;U=da$~8 zqJFw~`4_EUv`*edb6V$5cW$L~#w(s*eDgV<-qXHQ{B(J&m)|^596uZB&Cf3FkNT@k z?|L>r(pMh(MdvH7dFijT9@T4~k{`|aSFSwH9i6v(&$h2T?OXjz*SC8jPbtreQ-^%r zk6&JP`_UKOd*Gcy?*Y41XX|*i^?~ZP-aKmSOZ;Vv!>{pOXFc*G9o7Gp=3*Ti&A)V? zOL27f;G+BCzTjG?-o>7)e#9@oxJC2luPs0F_wTy-AU_-Fe#g=;TGu%BxHoL1u6=?-WRvAJv1^y%Y7*#mm2F{i1d9E}GLif4X~wbmxuqi{jl6q$!Zy$9YY^d7KF`PK84`a*T<$2@A=_o94_ z^S$%4)nU$Tq?h%}D~|swyPl2qS#00p=+57qkzVp&yzYVXM02r@zHB-+p1$%n-ud_M ze0BDs-?8$TbJ_L$&b8FXr8{4=Zqb}d@%F1t$NKHhkJZ`MFLqw$ORsI7bo--y7p-5k zul2R7%j2EB=-m}p+xv~aQ{N?BU8oL}cd_riIQdbXSY3Vf)5XibX#Jvf@-CW_b6PYf zbbs6{)dz!qk7~~w{db-#m3ieKm6{4`;2s?BOU2}$C~q^b^JYNw)wDe>GkWgG{@@lB0szN z>BwKYhvI6}OLH|ZcFB+Y_E|K4Tyb=9`m0^tI`8a7zt68?z-BbHJ zm+s&FthWxWN9**(R;PZD?!Nn7irxcu+5GY?c3*v=Ip~L8TmPlFrSo?m@#c0>-zXn; zy*|EDz1FeO{_YhU=~&%)Ty#Hru4|q8N^@TNTAk-ziloY-a8_1yiA6?aiR zi~3#uueyAtdyMk3^~tW^oWx;uwmgg0m*OwI{i%0R-tyWR{u_Pxt;KU4B##R`*WSPZux$qV{OnR)sGr*OuD|TpJmTbIm+GkgYRix2w$_V7`~Q{pMfV1sw|mbPkFDGH zP@Sz0>sKH0xKE9YZ=$@z(*E;iDY(M5Fo}Z2M#m3cEH}acHX-?JI z^*e9mfBhY+o<)5x|5shU(s`h~Z2hs9UmTj7b+zfG^`*E=*Pl8UE3^&%a7{8>fVX^>Eh*Iw0_Y#c^A#8=XTlV*f{sA`m)^z zbRKMM-leMx%@fsOA8|-;es-xYl)rX${)_8;x=tLOU8=wOt1UlX`NX04ztVZ3zc0?u zdAs*)=YZ-$b6{h2>lekVul-iH56ath?MGj9?}2v;y$9^le(L^8buY@t&u)Kgacrcc zy4m8)chNm5<+F~TZ6Ed5Ub?tN`?$wgy|y^#f6;wWC$4qox!8WpPdq;x?Yn3lf9=-G zZ$70tRcF`lypbRMj$P~3^OfdPUU|%^TzTl%dFg-AdVcnzxKjM3n}hlnCh6bn)^pTEA$Wyo=`4bGvMFY@GX5 zecA2RXJ2;7kNT`#o&VxGpRN-}XP4?<_ulekTul-iHKg!#6?MGj9@4RW6-&QE;(^oz|?TVCWhk86(n zonzNu{`EUnJzv@T$G-Bcd0LP1u=U5TzjG2_oh{ER5Jlgm+1Rs+aKLiT>0t!&)ewoqk6EqccOl}c=;EtU$jo%MRQu`Pj_y$oio;N zT`7L)y{CPr`04UEU-``w#qqO|-u&#+{;2=j^scYnxYs=5}l=7@N{mW+`J-e)L869<2AOI$NFg|4MzJy5%vC z+WIQ3Upl|K*^TF4IzNhI_g?AmE7fNm8||~`?^OMC^S$VvI0szo)VHV~G(Yk5+V)+v zj=y&6Ch6bn)^pTEA$Wyo=_v&W~V?VT>Z6D;958IEv=-vbG6nYQX zrTygjO8uffusXZ_U$%U>;`MXU`Jg#t*Q?unYA;>fqWyc1*=6J1gNyEmI&rNt=dyF> zrz0Kd7tLQ>ZTXPjT+sQkk&fnj`5RZ;?^yLLwyzcUs>@gE59MRG-{lvF`j)3Qy_Elz z`VsH@aqf%Ob$_-#kY3tX9NT+}-cvTx7sdH5(RazVzxvD@SAM#Gr$D;=s2;5Dov5EK zUj9Yv7p;?b(VU$3E1hHMK2@*n9F{H*y0_}9E}wHo^F(oY?QdWEmHeo_+SU27dD;Bd zBR`6(O~=KdQsI z(AoTS=kA=5j&yY*y>xD+xzI0LpW^gU>ZiK*8tZ3UZ$9Pn(_bmCc(ywHj`jQE-=h8( zy?-cf+4NGq*6V9goOvvoQz@SAo*=(_%&Uc|7Hp+|YU}NVhj?PB?p?Q?*P+!T9{I%(L?b|wPuk+R)+Rr&$bT0CXXP5G$ zzR-F5KG^0X&N-I)ptD=2{ZSltKl_T8-+jUM%hu0D?-Y6u*rojDezEc9g8u&L$NX!X z$HmUkz3jTBH?H}XU!1(o)w$JPdgItAKdQ%^>1=*_=h}JKuTB(Sx~HYN&@Wq`;?%7l ztX{iu^|R5O(0VqC=Vvc=9ou^M|KIu@>$_UtrG1y5F0T35W!G6>id*xdn^Rf8dxHG# zk$ZOp?CcZAJPXUo56{i1d9 zE}B!%zqWf)cAoCpx}W^LH@1f~5{?%)XEA@x=a~>C+pZg@9UCNKG(|P+o*ybbNIhOjMvs;(_Q5<$Z`-+#} zeZlt2*3ZTDUR7ty|4QdC-dxakrXTaKZ5|gpNBei((i_)&%P&q|=jz;QFTHVWlpob& zzH~M}y>son>sKd=FWuA9Ts^#b z|BK>L+_LFtAGBUyi{i{<(VR;0boT`L-6Qu7m#(fw{pi1Tb$R;!`n|Pu{!)CuyMCAI zT5R6h-g)HjzUo1Kws(ZhPiM=&X#Jvf@-CWF&%d^NQg)v1nR|$IbPwDQ6mS2F>bG8A zG*1-gyU-6C z=)8R&Z1WN497}!B*{#$5C=R=yeZ|Y~zF_-h>*u0(3cUyHQhsy4*m!e6fB*Dj{f(gUh(|sJLq}U)?f8y)8$9?n=_rwPq(i*BOU4XMS5xOrMb{AyWfL8O8r#t zeXO5tz4?^&tDBB%UH!Gy;r^GezhiwzC{O25TRieFn~wHD>-DuL&O8>)yEJdQdxiY& ztNV0OU5onBf9>k>^!@d|E}g#=-|tSJbagE@Z|%PK{PLrEu)24oe!6)17p-5kPToax zy6%B{B#zyC>YkO2a}Mrc5I652U$Igk(ufCEW z`D@cr{+`FH?cC5h{h|HDU33rR6VEQ?M}495_I5XHfy6uD2 zmuH0!_v9UTE^~d)63caIjq%Vr=ch~RI z{);O={eS;lIHY%9^&mgnJHqCtv*lm3e$hI47tN{XU)wz?J5Tq_y+ijBdvAQV{NB4# zU3B%kSMs8HqBwQ6AA0k%OZ%hxYSZFU+BC`esoTye(3DhrO#5leM)tUvk$gkwto8ESNBe#_kdl>|4R2jyt$z7SwG&N z+V0Oq^Al%(HZGkXn~#qCD6eyNZnc-*I5x_U>M&nAo1fmfcHZ@?6UCSAX=yI>%hsnj zb?XPK*KS<>Y&0jdo{i%9*^6Drw%+}}*x#`$-uvkLD(zdnwz%eFm)2XiC~nP*Zcb(W z?g{d{NA4XiU0tX@Hdbe&zS-Ve^d7U3z9_EWUB65FFRuJ_??2MJuX>Q5?HytB)7kPb zTEA$Wyo=`4^RMlml%1!0w(ci??~U))I{RN#zkAiZ=85+4UFe66@}fG}*tv?Mvr&I& z9;G_eSMno&Z93Yww14&5;!6FY{hY@|=jT3&XP5G0>vZ0}54QP;caEh#=Ie-ww^ z&%Wa2cVDpmvh{P(JB8i@b}9cWoxgZ~wDeO)?# zDZbyGKI!UOY~I>^@A>6N^VwX9f7Bx%ibLn2E)*}n`-1Y9`nkB?tLkj~IERbQN1VBo z`Z537=5eueRBzWUy>V>uNMHLoSLarH>5XHfd8z}eFS~Q?yz5sdiZ9*M(tPNbtxs|7 zZ}oM>z4{x+cHZ_!{)>%g`+M}2{T*w)@5=X9+P8Xbam~jrt+x)Vv#)uXcWE8nJwbl= z$i2JRerl`RoNHINzVEN!TTACJ#rM1Gcd4$$=B@3WNB-`s9^_|xN7(#yw)~6MFIp$> zqB&jnz&#SjcAwEba}SY@&cS_fPsF>YrTyvZm!~vO>^}Oy?k~UnN`6#dZMytu9;JG$ zEBTSXHXZF-+P`{jab@RXKj(bWxw%K;*`+?Qbvo~|@y@B#51rk*^jV5;U2JuTvk$gk zwto8ESNBe#_kdl>|4QdC-dxakrXTN5ZS%NjUgGS}#-;O@;*o#t*Ym2azv|1To1;3g z`m)W}JwQ6r)rs`7_nK`!ueLs|(;u$B>Tg{A>_zLHGt$eIhc3Rh`(JwJF1^2Ft@C}^ zx8z?GU%&jNIP0)F+Z-0n8?U(D7x&1$yV!o%>c#3kpT_n5^}a5hzZBo^uHR+nt{!aO z+I{c&yRUkXpY0uC^V8Y#FIvB7oxF?Ybln5@s5C!xAKg!+qx;f(BHleM)k#;se5HA! z{p{0z=*`bA?T_lKO_v|dqg0o5B|q}lrlWm-W$R>jF7|UT(D~G^&hNV|^@;7rd6$iM zPNjb6?AE2vQhe)Tt3#Z9u>G?2)9=3edatUp+n4l_-_ z_t)>OrSq5K``z`swEtrB*7nXLfA>`n^0U1oY<@ah{zdB-t&?}roUVJ|9+jP^`{;fm z9o;AQ2F1Ilr8?>Am#;KW>^|+M`^#^?k{{Jqn=U_^N2xCBN`B<8O-K9w%GSy5T76Prv)>-YN7Ru*>#kzeVTH zZ$0|X^yB@hZ4MXBS6t^;9{kjdWZ(e}Bi)OL@>7#nq;xbtqn6 zi{i{@(Y#CZqq|qg@4mWExO8=u?Sr3g9@W`>f4#3u=P$+gyVDQdJk?h=Z|%PK{PLrE zu)24oe!6)17p-5kPToaxat~hVUYRSpXC*(Hzx#mh?`s|E@Au~ZnJ3ciqYpO9i|S!x z=PHiQM)|RGV)Lu7)$aMYAJ}`d?CRE|xu81P=0!(3iog7g zcRs$4#m1?B#q;|eYYu;t?T>i z_tw(+OY!~g^hLKnsso$1ws#)+ldw)chQ`j_bc5u^F+GuANkST z(LHgWU+YkRzc;pdBHcdfWuv^P9yWHa;^^!}` z+SU2>SIUp#(Ruqm*y~uTHZMBTQT*j^yz}v0EH+O4;`xzZ|NqwSSa}y$J?cVn=CSC$ zp!qJl_r^VTZx^joAF7j$bak`&>FmD0-q)q`m*V@~>5FcER0lS1?Y{T?@}qjNy7|{n z7cc*!^^4ZYyJ$|%`<3pSc_Q8SkNjxvwcY2}I@Is`bMM&ZiTw6aFB|1W^{}yX6-Q?; z+P_psbv8d%XP3^y9MHVji{jAvwLkmX55?E6&ab~xeiV<++xNj%r#RGHcTC_h@KpNs3gs?L^Qf3I{;8t?nokNMX&hl`zOa|xr>OyhmvFN^_`7XQn#yxg#7p+qts*{a$b+h^D?7qK# zZ!Mj_6yNVoUv&GUI5V*Wuv^P9yWHa;^^!}`=@7=hy!1Yd;iU zyE?!AO8HSdI&a?xTb<&ZV=0c#_U^04{wNOV_Cvb-?hDF~*6HV>cM81+>{5Q``jzH_ zzHj}Qe{FNRXinnHkBxM5TfE}<(R}n_{v&{wV$2KoI(oy{7Z@lyIuKEtxNXMn~>remX`W=^weL?dr`Fn5NWA}E^y4I`CrK?*U(yjBJqW6@I^hI&~?({{sKdJ*)etQ4^)zIZf z^>1{hx+@yvCR|d_R$9$5#q*>4 z*z>Bbzv|1T%a7{!UC`P5bo-h!(vfaoq?hJinhX81eGlT?r?T%$ZTD}jGZ(ZU+q~#V zNAXC<>TKtuzSdv8es-zvi~SvIzeW4-%g^>5>uXV*`7D}uX?}F~3i;hv_X(G-uCjgb z)6JziyYH{xTTACJ#rM0@58XV~S2k~L?>zF$kLtnd-jVw0;^kkoe$hI47tP5%c%^$~ zuIQeX{Am8}1G>Mjb*R7JoBL;;NVktZ*eEZmhmD=9I652U$Igk(ufCEW`D@creseH) zc6s@or~c4>&iSH!5#q*>2^t@{8ullm-@}v6u9kBVu+t-|t zj&%DXy)^gIT@QhW1ANp=_nrQSe@;BS{EDDi%aL%_bdI5 zm3L8oe)-wnWqmD*GoMBCF3pebULn8x>OSGp)m63+e!96-XZQW}zAl}=6yNVoKXmg{ zU)j91``+`*kLtnd-jVw0;^kkoe$hI47tP5%c%^$~uIQeX{Am8}1G>Mjb*R7JoBL;; zNVktZ*eEZmhmD=9I652U$Igk(ufCEW`D@creseH)c6s@or~c4>&iSH!<_yXbqv#=A%QG5^}`>qYZg zas1}%{w(VE@{4mnovU-Jy>xMl@}s)Uk|+UmJzoj8567xjtOE&3fR{^H8B>Y-bwuSNIemF|sr_t?E% zv`&5MVPo~$jpO%@Vc)mv{H3^lclx5+7uAK$Tf2WR_~l3SV0H7apDtehMe7%>lXuaa zocAl;C*Ltz&qi}!bPnDNc`mBo_fwiD^4murY?K$(!N$&29G#8UA-zU&XsRJZxk+5B{Q?T>V%yLU)0 zJ7>1JyxKh<-%+W*>V9|8#Ub5XP@Qb^q9Yx}U;f6op4#f9DpRm*0Is`O!N4T=Y(%_kdl>|Mfmx@%l0U+U9i8oK_saxtiPJ6~~X}qYv}1 zy>xMl@}s)@U9kDZ%WHq6Bi%V5y`Sn-IkK)mJ`##v}6z3dEadftMsmJ~(4(awoy8P}7%8%CR=c0ECy$9@4e&5GM zb69b{Z~d5mZT(#|uNB9y5BF*DisMK9^t@{8uR42CepI*l(%JlUdF_vMq`RL;FW0=y zh5xd9KKgKvuzKyr)o(7SF1C5mk&faof8$$^ddtRnPsQ;gU0lgut~kGA>#r@JIu^y5 z$D;d!=DY0P8~50~U9_(Cs&nb;ZXDe@?-_cJ*+^d$*Y8eWbo-+^aOJ0a|B)^~st2o^ zfBkgv@-JGyXq~)^=H$Fz>AsmK(tZEPkLF(6eSWP&{k}W*j%}XEZy)utQC?IJ8#`BV zboQeCOLbId^J8^(`FbACrF3q{-~Q}tKNMfPI=}u(`B6MNZ{G)7o#LEhDUQxw_eUM# zkZwPu%kREm`?XF#7rj&HJz$seJJ+u?7qnhK=3mnE{kZbez5hs;AJv1^&A)!Sc=;EtU$jo%MRRiAuXLY$$7nqp&AqJOdm+z7 z_4|HG^F-%uAAPV@fB75VdemDs&fl%Z)qah$PCPEXzhfJ(-bH!MW6^y<^IdlDjeG3gE?U=m z)wy(ai$l6~-c$6RvXQrj8cH@0~q-9GANqr9jdHg>M!=iqgERb&7M2r8qjf_tgF<4(awoy8P}7wqNV?b8)>_ z)!Fhp*RM1ev|c~vU)vlmcAn7}`s=0d-0_gtL&QR^&9ha`J<+mQiAstsf z@%y6zFsshapuTI zx_%a~IDT}0dtSBmSDn2mKdRfD>1=+wy!J;r(%ny_m*!fU3;nX4pE!M#`l;?d(ZwO% zT+n`O@pPo4_{-mTb@V%6qkeGduYUWAFPope=y$9<@-K=rk45(d&9~(5y>XA-+ePbI zuR52mZgEJr&U=R5V>Z$k#r3<>7v27-4qW-^eh(sDepC-uH~;$S;^kkoe$hI47tP6e zzta6OPjnA`|HzN#j_!%|uXU*3cPI~P2u@pyVyI0oRAH^YEJksTNUr=1BpNrlp z^d7KF`#6V-&QqMZl=?CM+UD>p-FNe5t820KH;y0mWgq8Od+Fj90U`FY=>xNJl!-kuJYE^JDee;!63e|7y!ust4s~yHD); z{f-rPQJzKfEZvVwcaM92Ys;sO_OVP#)x*ZlRUDm- zbd;xTedxR|5Du2`OSwd-ki}rWvi22@}u?kS^X^Cx<&a> zJ>D}q^0Vc&KhlwIPDn5HkM74;+d6f#7q9%z9mRj8{ndqZq+fI%R~_?)&Ta*3$V)@%`@fMOUYKoEtW8ZSOqt z%a7{8>gHcRUA+8@)-PHo@1i+5?^pKzvF&57i_XD1Hon#&&iQsf^F(p>S^Jv{+J}w$ zMRB#O^DnkObs|4jXP5d{&sR2|c{h)^OBaWYH|J73>I?}IJBxaMVxt8G7iHqw!9 zPWGkC-+N#^KU+T+y;JBtV3+d$O8u1jGXL88E5$9HU%hPc=Dp|~_?J!hci29^Qhi;o z4(z^kTBukf7yM1y{}8>FU9w}(-(c!YwpckyYD@} z{HPwRZvOSt#mm2F{i1d9E}E0`er4|;+dk&H=w4XI##M*-`kin0Gfxy}AFST|YhUw2 zakc5@b5UK^u@@W1F7?^|UhU>H@2(ej>Ef{Q=2&(P`a`>uYLZr!54_x8c+ zwZ%FAi|&m&ajollnm4^R9r@8*`4<~kTR!A>UU=#J=7@Bpqx`jf@(d?>wEid7?P9 zAJUtjUD_YjSDS9$D6SN5KlWnd*rh(JzuMj3yjs6}jkiCYjdYZ!6p!ja=k4CJ<*}}H zzn+Ww7Uktvmwiz^;*pN>v>$!Zy$9YY^d7KF`R(_W)}#LP)48($O8NEA_Wdt95B_D- zy(|A#z0K1)>|4FIIDI!S9r=5{^wK$&_N8C8dm)a!C@+8Q*3XaDAsy*RN4kB~&5zY< zi!0@?{;MrtsUDP{-TPF(dtX|&e#bVRZJwq3dFk$__(gg2*}j(EICkm1!M?-Q`Ac#7 zLiMw|zxuIxYx}!^{PLrEu)2BIPZux$qV5K zXg^dBinAY9mlxH+#x)nZb=dP#SJ`^hiTqfdU7CyfzfvFQJlSa9=I5ugk&fnA%D-ry z=FOJpnnUYhqx!l(KZ>&-TF0-hMREGM=$%6E0lSp{SL(OaPv_6JpSa5whb!KEvENDc z`t{K`b(??frHfm1FHk+cFFKo_ZeQn&bfmWqdT9=5-d}C&%=4oD#MQPx@|RuDPe*e^ zI?_>|+C3j}rM%UDwdE_-hw`)4%U=G*)$Z?D`AX;YO7}**d+YmOw9fwRi(WQv*?m9# z`@quqOY!~g^hsBzddwf2w|4)%f?s}A4_5b%)K3>L|DyGa*2%kQPR{$4?o+=P_sN_u zS}%?*PpQ7@&ewUOd7?P`=mY8Un-9C>NA=aN&fk1&^;lQ(BY$l=>QCHPT93|6-Dtn& zw~o$6I?7Y>qdL%em-5)pIh6XKv#qnQe5H7C>gsxN_QB?pM?d}UtFQN}I$QoTtpW+wg zUwzTb_P=cJ4SIjsNM97EFH}EUo$7IJxboB2-wXLsJy_km>!*vi@1pgK*2%kQP8Yl0 z9J@|E?lbltcRl+n?eBb@Cz>aUE5+00MRl=p&4q3qn$N}7qb}sf>g>{asP8NFjm}kn zXy4|yj?P9pI=52(Me*j$me(BYTk3;eTb}MO4z0HzTE{P6*}C;}alKd7+1>Z6t&dWk z_Q77%U;T8HpDo^;7kxMU%ci^M^8ZTpc75xx@6!2k<)x#2drtJyJuU4^zii)wxb9zD zJiWGlkpC;?*9X#(eo=o{9r9jruYUQ;)+--fz3lqkKldN2*IvJ4omcC8)!k3=i}J6& z=wEh*Iw0_Y#c^A#;V%M8v z*QsaSb9&daztaBB*Lk9OqPS8#U0zfd8`oUu)}i@aY(45ieyq+ey))|jN_}I`Q=j&2 ze(UILq@!~y$F=(XkP{^HPj`=NFG@|CSyKNr1I=sjSU-S?}lk5ZoY z!Cur~{dAO{E#92bJ!SiD*d;$&FaNJpZ`a%3ezb1r7cdYYjov*t4DSlD@ z)fc^N|I7B?p!b!H^hI&{LiMxNsUGKsD?h#OJzaiO4^}tt`sw23U$lPFI(Zk(>0;NL zW7nz2{ayFE>)Bsvf9LBw(LAwn%a#|_#l|%kx^-wi7uB!6k{|hN)6sdT?<@6<&RKtG zzvj1&&PF;qw^II6|KiP=Ew4G)x6}u{wmjWm99nNbw2oiCvUTg{;(D*Dv%BwCTOXx7 z?Ss9jzxwGYKU=&xFS;lE%ci^M^8ZTpc75xx@6!2k<)x#2drtIHespiY+Sci(``7OG zQ(HgC|CRPtAJUP2(Oj-NaUE5*~>m%6I2xzMdc^SRi1)P?+5on3lo)c2L`-}&kf?c4m; z)7eNz=T^$UDBhge@|uHvOMTF5%hUbEq4nluU+ehgD_ggIE_$cXd%!Nc?^jzNr9AC} zZ62k2y>xzaWs7eewVh9?-s;vnw^z!u;?=$4mfn2kseYuF{C{Qhnp5-A#b5WgakbrN zr~X)dq2IP3ITnrC(Eke}WBbaem3VRg23i+;zN`&XLtV$WM$ zi_TA9wX4hH-d*(0imUD1u;)Wp*P^_O{`XL&y0Gi3t4AC^o!xzETfZn?-bM3rADn+} zI`S{B`%8BY?t^GI2mtwTP#^YOpSLhk{)l;8S` z?u|HgpgzzXYwNRA&(itLmo0wnS6w|=f6s^gN_khjx|g56;^^k?9w5Es|0|o<{F;|8 z-o06DTy6Ic`OU*UDC_V2so#2f>3qbo^;??Dt1iwueU|1}-8$rFH$NTS-$m>B*^B1= z^>?iKcK)@+p}H1(UbXc>m&d)k==Yhp+RhDoK6Lvp%Dd?OFV%%zUtK-o`04EKQ``DQ z@$xR3mwVvcYtxZ`(Ycqt3+LcIh;zT$i(Myg_4X}a_qUFYbfhC4tFzafmfyLkUw))t zl&5shzOwt7d-JpP_sTV=>hhTfKi&SV&;6obbYH~D>pQSt>kzm06W{f8`SpdZ!#;H9 z)8D7wgX(Pgt-t8rh*JmZ1I@9vK1=n`*`@i?#jpLUs|V{}=Og}=&BK1B{u)nT=Rhx= z(?$1cvFpuGJb!I^sm|(MFAl{!e{?_CNJn!-I?@;Abw1`$>f_RTudG|te<|L6wdq)Y z>!@uVoxNzDc>RtwSMlbH>pbY<7R9N*c6E8y@2=|JU3zWrDE2}E)TYEdiPzryeJOo zSe}e{<>Zs^Iv)Nhx#rrf8)&+`Q=fMI@#{2y7?DZoILh-F5Q3K zXY0^R=egDrte;Ckvz%K14-&fin+m|`&k8M8ozie^1;^pst)%mgh_ThZ# zwU^#FHd?p1;^@xH{E=Sr|CP;ae$7j7{Iad*XCocyz7zKY=}1R9(%mm}UbK#XouhkJ z|I)?r>vPeZ7p>!$w>EuIUDnm6m*!dacdYrNc=I)HT)Mgzt>b6wyE?z$X&3!I6Ia`N zjlIuw`!346=>0F%gvN3?Cw+B`bF{bE}ECQT=(JfyKlX>Y~PRl&^@`<%j15r zQ6B#GO>bQFzxta`pXOKU8=Z^#KdCF_u@;8t3tj=zqm#zPcJs+rRb6OTXw| ziDP?D?dM$7VgI$C_^zkRZ#}9*exy6M{yy~{RA)E8e2d-JT+kf!L$7V$(mkQGOLL-& zx6hg{z3e)DH15@IzOIx1qW&63?|C>kddYv$xh;0Re#P_GrkCoh-u2>8yz@u*gN<}F zN2DWtQC{a`4y8UWz4yerMg5oJ?N^(Q^|y}N*3sFE=DFy1ta+k%^Ih!utT?)O_1CU$ zo!@B}{XP>{+k1_@&vbR6I#Aw4?|-Q-?E32J5ywwwcc0qUFN&9U(Y(x|^I@0b-7gfk z?yv7doILJ{dn&GMUVgef*uLrAcj@w?IHY5Bw*B;tOZN^rANi5~mCgI>dU?!$<5OeZ`jLENAXB6otHRv`)8|@jn&Ny`4`tbNmaqF;M@Kr+k&e~boe#TywsTRx{7Anj zPwAe0W%o1p=4Y!H^;cg0#_3A&K7S!G#_kTjpIjt`@GV*t$6j?XHlN|>H1o9 zPq1caML~nQmU@j&u~i{Nhl&^I3D#fBj1r$FI+_bKT`eD{EO>8 ztvcxPcR#v1+Xr2qvUSR*PPXs=mFjS=7u_3i?iH#>Ubg_SCI7|Vx7uB=Uvd1k>7_cWcfB|i?>?aKk6rScBhrz+D6jMB zIkU}&jZ3fJ97_FHXUkukzG(kN>-gD==4oE+k{|t!HQ&Xa&x)gqSAXs5)_MOfdVj^$ z_I_gTGhJP%4wQGX|6ajX7j}Jh^@!uAv*lm3e$hI47tPBYIv;i^-u*&x>;C$F#L45H z$Rn<7UVgef*uLrAcj@w?IHY5Bw*5+TD7{C{M}DM#W%K^JULNybdD_=yJO7L3it?JT zdeq65SKa)JjbnRn?eAQ=fA_QAIC2|0_|pERbxSvg zMd$rW_sd+-J?nhg_OGpe>*ZVYUhtcf{PMBwN4LJTKVE)u7tQO6<5%~pmwr(m6z|>t zO7A4fyXbt-In=JsZ~nE-iN6$&^1I*ar7w29@5_Fz!#T@en{GY7Jjh?F2h}Z(tJfJ^4!CF}jE5T$(qXtuFVJU*6i)_1}E*R<{qzr~c~fuaxgAuX^?M zE4}+)+56QxTaUS7&ykJd(7bD#Cx0m)s?+^u*H6d#eUIIbtuAw5%ZH6;UsMmuXFlwn z(`8$K#qqm0Y;!?%=wJT!=R7X@yCvU6`=NQEI{2}EarCRskSHy2cg`N-eC_%E8{T3_9MXkI9uAL|!KSI1X+Kdob< z{Y(4O+3Kp!Upf!0|20q7$%iWsU7UNf=pOOccK?vyTo%>Q`$osEr!R^(AAgU&()(QM zZ_%8}?r%NaKJ5CPSEv!y`NBt;{Z7xU`M@RF+rSqdWl;3`I z_l{mVx5eHIaq2)iR##8`^u~)rb+EDP#L2s8UcC>O?OvHHx@Vm)+kUmx=RV7~Xdiz4 z%P${Rx4yJL{_3|*9`mAKyz150MR`!XcmFHhE0lN9IbXECan6^{HrF*r@hHD@VDm59 zk6+*P*3ta6?W^vtr^{ce1LYIPR_80N|9T(9cOUmmfAZ@O&F7-`NItZVjrK$HLi_V$ z{o?5Md!>3-yuR#T+Lz8&S9Si>-1T6PW;wk_2w6c&a2dC^~>&eV7{d|-v>W?vA<(iJ?cksY;!@n zI69gaE}b96q5Md9&(u*m*G2aOS6p><(w9v~@umGs>y~Z~i@gu-4=%sCqI=f)vYkt9 zb#=YG;@L>ozx?uHb?*(zXaDN#uaxgAuX^?ME4}+)+56QxTaUS-^Qv8)-@I#^Cx0m) zs?$AY*H6d#eUIIbtuAx${m6%nXJ1qg%4a_8p3`Mpf5q{;H*9l3b(n|z?T`PWIj;59 zYug{~x8nHGI(5)%FTHVWY#nU-miD2Ww|jx~lK-Ooi(TKk`K@QWCyU;5{@U&x^4ot= z9p;C0w7z~iHlFSKFV*>~i?eRgoJ#A_dbWGPuHSi;`mBE0?mL=q$&dVO@A#tMvHB}p zCqJESF4fKHqVo_}TYj{!{pi=cKD!%#HD9-wommlqe^7%fhvrFrcpDiEKFSZYP(cFHe`|*{%Z>_WS@W0Y| zUNlehDAi$Kw);%C54)^i|IV-bs>^=t=IuK9O8cWc@*~~-_WKf@GrN>u+*g_ls!Koe z*QQ_8--_e^nqNH1FCUxVe7)1Pm)dQ9wo*Nzc)rs`dIhN)_zifSqGuK7^ z^VgOK`OOFE=zCldw)chQ`# zd*B|G=7;WK>7L>`2Y&aoR3}~i_HSPEL~-hCKlJ8jm-a{X)uylg?2F>8EBTSXHXWN^ z{41@;bzb(n-b3esjhDY{z3M^dU0Ns3IhD?t&Tbw0EXCWWREIeGVD~W>{q(!9?wvyK z0lSpHbuL@|{HtF5czDo?jj8iz|=wv|e1#tG52CFPpCJ(z(#t>S6c%-48m_ zod?oO^_AvBzie|6r~b0@=GQ-49Ma7LtwZk&8^!ap`4=0<{+0GG=fBCzf{mSOG zj?PxcE3Fq_+xgOa-`uA~zX$&MJJxx~=RHO5DI4jF;`-h7yR`q}%1{6AKLbE|_f-$_ zv%Mp1emYzJMe7%>lXuaa*7?)jBcykp?x8rOqkDkzxToH`(*32Y-}=%#kzZZp5MJ>qj-Kc z|6=3V*nHLNXP4G3UiHamJ-b|fI$BqI)j@Y&SiN@di+kkWUG)0`uivry5|7orx9B}) zBYja^zq@{y_Fr82>E3^&cVG1&KifOP=BKmeU$lPFI(Zk(>ADB*QQ3LAm!*4&>t67? zr=>dS>Q{%nXr4${U;CjqKfAO)s;@SE?Pp&UXI;sU{I%&Q|F3jDW#?i)?;AR&+SU2H z4_lwuew=sNc;{5=ht9S?>a!Hzy4dOvXCIUw%|k!^?yGyJ(0jlxWcM z)E|Fsd63^c%&V-w>z2QH*p06(?*FXaOR_9Ea%IsOu*msOwmPvoTjYqs{Jz$fqd%wu zFkFwwO!Ae@r}epLpUSQ)pIv?W@{4O8Eq}57iDR#>-F-vnzv%CM`T9H7T=MC@zuwnn z=P$+gyX$vp{l(U;?t9PQb@hY%w0DH&XQ$O)l)orX-9`Jf&Y#^qqS1MxxYE5u@g=|a zu5@qN_3yiKzUtafWY?cLkiGTk()y^sYIf_O^-6wuB|q|4v!ivyeWiR{=Vd*9bWV$n zqs#VPzw<8jjn1ib&g`^2>zT6@FHRonUz~MN9@;PS^t-Qpy;rr<>g($(&4ucj$NN() z{-S;wXMgzV#w{%##i>X0JFl;FKjhJ9ztDO#dcW#-{ukXF=YcCvf7t%3#UVd-U+ph0 zzj+(ic>9C)4gGzfQ9M7*zt}ihzw&5g$7QcyTxnlQajip_jlcZjJ2(49^~F`2huytG z=fCJY7yXXazc?Ce_xlRHqcpNFit}Bf?~=BDslQd9z5jU|yZWdftlc|NKf8GK7v(R? zQ+Lrmt@CGhkC5H|A^%0^AP?PB?;ZNP=ib&XPaU+M$gV$gAbac6rM^&q)$Gy_fI zM=v&xF0K13+b3;aw61eO`(Ir_ny|TeEp*DfiC6iZ=HuceaW*f+9&ZS z53SdE%*F0KSnpNswEFt`O7o$5=IOrDuj|-n@z$Y{{gw6?<+UEadA_peOr!m>4%WW3 z{d5nI9ogGIduiWG`@?=|bBZ%xY5!~Q+~QDtsW z(md>R`!HuIzJ1a95N93iI<$HE-Pi7&Lhk`xs;|GVl#kBSJl>yb@fY>8;`rSg`+HFy zig%tgKf80QUiQY(xaKksKdm3y`CoKzoCmHv^OX8#zqE4~XI~f1$6u`u^1FBTrL4dE z%8$(#kNnl_$dBx}?E2t0XW4z?mxr}?eQ~(_wbSY^s*CcLR!5xkHt$7$@6mqJ$ll+v z?BdXSir!Ni*%!t6F41>MTNm|(t3JDbPl4>}qkgb#-v{)=sN$el%|>UVh1s{MGDe zK5<_uAI)KYw4VJ&=d;*2x@_P2K=8kM^tc zn2X(e;GIJ60bQ!Ezl-`8XJ1g?=IK7t@-KG1#`9bMmCgs{sZaA?^O$SJv6uGCI#~PC z)^!h%9of4+dud-w`@?=|_gtL$%I-IR=N5MzP)l&9{ZeRAHfbl>bNvb%5C`)uFQJrRH1OMd;!FYPDt zTSvdh-uiTDeYAho?4@-V8&5aBcDgj*uk1O}-52ZfqkB?aJOApl^Eb|Um--U#oJ#Ys z)B2RJUMXIm;*njPbx<7I5A*c9uiZO^-UGT+|CQ~_zM${gJl>yb`*6|zinBhA%g$el zNB*^5`>$SeHjYO7t`Dqz>7Kv)!H(?uME27CEbR~brOheMe7NSSzj4dI;@nSpD4w6@ zUu+!h{(q&vPo=)lIv2&GcpCM&v^*47z4nRS_l4{wfA5QXMzP)l&9{ZeY);}dsNy#bPwG> zWJl+K?hC(rTI!Qs|LTXx>s^@=AW>uVzQ<{>t`AcVDdM zUR-o;eGlo#kVh7AL6WoU57SLzx&$PdsREFzVrM_`RF{& zU5i$i|2 zuNUpRxN7y0-+rKdL+?M0;`wR*#m3Rj=PTDb>X@_aKJm-L+O5Z5@|T^nx^YV@~e;f!P>nO^|Omte^LIT zJarfC(>i~4_XydYFY;d$?_QSPJM?$ny{%n->)KBgr$2Kbd+XDszEFSF?AAf+mEx^O zFE)-Yt-I=6cKN7p^P+XVOKATW8%LMc$NHVOdr#|Eo_^8yK$r6MzxGR>zT{aK?UQ(v zht}&n=3@6Ac&E^NK$q(4>!P{EnGf}Cp6)9x{$kfrzj^H9FE(HNny=?oZQk0KX4iN7 zseNhZ?>-&3h)erL1-VvIgomPKQ{-Qi}7wywJe|Gl>+3hE~XX4O(ME9cP_uiH6 zFT4Jor#{erBD=oKf$XhMm)1x9RkNE9?N_NEc_lycSF@w}ex<(4?u+&8_eJ-h=SY|4 z#LnZq%f>mU(md?6bJCAFOY!!p>*`yabx{4%JpJx#U+-1zv~~1*(LCbp3p#J}cz>2= zU$jr^pm-X^m(Hnnn!n_CZm+aY;;Y%sQJSaJ54(LYohy5_^J6dd`<3QWUp~9PPmAt1 zf3!Nt;(8zvKA1L3u|CP=Q z?VI}U9qn9EUUlvK&a2wG=-kk~M)!c!Nt;(8zvKA1L3u|CP=Q?VI}U9qn9EUUlvK&a2wG=-kk~M)!cxZ_#rF+3%-SZHK`bGDZz9=5e>)tH7NBq_Q``-@--Q<(K+Xr~X$h|COum%5R@5u6FNK?}2v)`Q0ma`-kqec$7!G5570rJ@ww9^YmWU z-u3m-eH5pUMfI`kvX_mcQC}DBuRI#rFSbAV-u(H!W0VedJu zUwQg1_0LY5+xe-3&O@Glx^Cmum#2RB!5rv*`;KY;S6b)meGqS7N^|MYeCT_>Xdmqh z%A?VKm-fAOn!mjKji=p*rP<9_ny=IkyZtSlKYO)v;4huuSDIUW`Rvzw;6940?mgpo z{%C)2+3R0%>wMVN(a&P%T5+#-^-6P~dbED&`ps3!>pi3wR~_|m#jU!_zUEr?&QF(q zZm&c2l9 z(x3Uz{#(t@AK9yW|M-13Xn%0o>tAu}eAv~|&tm6Vaj$muN^_ukw0`OO%~i_lJ){>`9raM0 zI_&)}mzKBYTK3LQm;HTM{p;^o-(SDiuCIUe?k}o~E3S5PHIK&TEiJz6I{d2+yEx~K zy%+4r-+f*8Y~}gh)REtJuiAUa?i}T#IZ<5uF%P@?=-s8!`Y5iNU0)Z)TZcyWi~2zM zeg9vz^FsTkzVD89t|+g%c7EqoZC!M3=w74q=eG~~)Hkx1`qY>8)kAf>!z<4DbshD~ z&SNh2^$zea+TYiE&^Y^2n#+DvU*7@0ee8bl*Dk(v{xOG14>gZ>&bFH{nyLzQL zP(50|bp7Tk<@FxYi>r=$C{7*rewRzjTXQXY=ch~W?8WPMtoIvzf4)z;>+2t{KGa9? z^u_geW7)5I{POifzta9VZ|uEbNB-{Xx@RlT_oj~gzI)Z)Lw4sVAG;6Y+K+kI)kp6x zjn+qT)$IDZDBe0WvR~8(%J=SnrSn4jroQivcCIL|x^{l&Rc&2#Zs=a4d%$lW^r>%T zFZHP}>#K+Ac!yV<^Xoe5mz~F4?0pybFWTR)_d&dUDa~a+s_oB3=hp8=UhU#b=TJM% zU-CQOS2|zu)$HadI~T1V+WD8xoxR$*@R#~U=SN=@kLK+?sdgXPtKB=~w;zkXr?t=g z^2JpzyEy$UcCHooYFDo`52{D&m#*KOrM%ul`jxFqH=n)V<)znL%ij6v((kQTzW$E& zeflox)i3+ihrbKIQXj8e>#ThJFaNT8uX+!>Gsy2=vD-g%uf?N0+I{f7(eA1D4xOj> zvi7d8kM5&5eJrYvU6;LV9F6+AXn*C=$bPZ?$@lJmrSn4jXC3E6JJ(k_pNr-|_Y2p3 z*RMSNmilL>&F%cuLFXY)KV7%+>dRBV`_Q@EZ{IQ9xJz5->wOS!UrKZ7&wS{6zi1!r z3(BKWeCZr&r};~M=k`kXL3}m4IZE@C`eC=9W%s?h=fjWs{YrDGFQ2{lq}qGIUhUo? zzk7`K2baD66}QfZT^;=_cCHooYFDo`2dYQwm#*JjrM%uldU4fJ55=j&-tTg0d26m^ z@BDP>_tsbTcdYjseRsZNy6fv7r50vlS|4QeD_Dy}? z9qn9EUUlvK&a2wG=-kk~M)!c`%;?AepK6^i_Wdzjl9~$m(HPfn!n_CzOQt?;;Y%sQFbm`KeY2Noj-fEbKx)b zi_VX}C?3t*ds6K_vRAuz$ZtOueNStj`Q?kNUUqT%S?pXZ?$xedX&zLM)-PSZIZJuH zhx993mu^0LzspOnxt6{2)1}{AzjFPK^?mv->D4d$)rY?ezfvEsT_pCb#b@cX{bG*Rpqhy7YVNqVG*y_4*y_yA;3rWf!-&zHjl%uW$Yp z&yM0tbCu2GXJ1sm_rN=Y{O%RI{X_R!Jj$cp2j3g*o_g=lJl@ONyS_f`KZ?`GqWai% z*~`Y!sIQCmS00V*7u%owzW=YnoklMRTD0h3mfSSDt=L{j<~Nc7E!h z^N^>XuG@I^<*DC&=v?l#@0eESmDYK^2aU5YrMdLiethp2?W27`c{JMZ(!SSD^Ou*u z@wEG}G`smq^OgExx4)(P!d~q>_)F*amF8AoKKu0^xR2thd(ZftKiVH$_WD=cIv;j* z^t0HxR@|#yz0w@09<5)xesh)bdJpNvRYyHsajWjKuep}J^V6kw7JXMVvM-ABJ64?U zknZ~WM{$elmEvos^+z|J#>UG-??EZQaqQ}r{5^l~D)PHm>}y~7Q9dot_eQ(F=JW1V z*X|tUW9Jpue)PqzK6-cQ#m3QS{fpwQLnHe|`-bxS{=aJHh4xMT-am1#bPq0C7o8ir z*CoGwXy5IhAI;(Xtgjxb)9;19bdJ_hzx&X6+;8jA>b%lAulJyF_N6qJ{@RcIx#-+{ z*C>xh@ul;so#rq3dv2FDzqo34bCl*O^~3I7IR9#P}BI<)YnD(E00F@i|tQ--~U(bywLtx z$9d4s^_9-&qB+q0LiZZo1O3X=?<>vi{M13`p`L!aZsYBjJoUQ|<|*Z|>*JNydA$dX zvoEE&^rcUzA(IT&29;Lwa%5Q4d#K z?OlhTea*G(ou4k??L12OKIKqr%&r%v>&Z6uXf*m={#zu`OC}Sc-lTM&2GNZe5HQa-M7Ww@5VWw z+MORdH~OM@v>*NFoND)ry}JMW#NT^he~_Op`LX_WKJ4Q4v)H*-+^b!^(j2HBtzWu+ zbCvRX59!5KM?G9|tM0O|xt6{2)1~j_qI)l{y5F1cl-=*xuCISocd_58IP2)Itbb{C z{j2}Vt}9QyMeFwby{pLYUa_xz!Nc*_gX(RzkT3;r8%6R_0>c5`n~X% z&e6W9-+fRQ%~h=)s!NyFdA$dXvoED}^rsKx|CQ?4Px`f2_o96lSFJwsd++T_S-<@jclpJmylUt3mFmc&&9`Wu zN_9{^tvv`_2&+1(@BJ#wDNe^I>qQF`yNzXNL5U+daW zw2uDFf$ZKz>mWO_W9_v1rS(dAt;=3EjxMd={8ud>_3eDny3Pgd|6=3l()w7x`|f_z z;^gTU#nYuc{deE^<%zRT{3zdkp*~R^{h0^Z-ShQc)lRFgpRY6z>f1c_qgvcW`?uow z&7qG)@nz%dw=d?XUiQY(xaKlXDW2UuxafYkFSzo|Q|g0#g4Vp z*tq`g@_TF9`AhM>OY~jR*7rMBU)Z|U-g)F#AN7N^dnf8=7q9-J{6%@{F50K-9=Jy+ z-oB!H$B**RzPp#D_YVF2bAM}>-@5h_#py$T$lm&NsV~%DHM{y~ze@4eqZb=Tm*!k` zF1vg*hk4Pu_PcaWxcs%#*2nIj^LFoP{mRoX`X1;~zCPA|$>@(_jX?7Hk_VuED7gw!5^4lluKGn`|Uy&WZh`luhQ-8)e~yLj~%k(YvpH?ed+c`t}pW z>8taww?19!3-wpc-hB2&@ypNNIJz{iyswmx_RG9zU4HET({m{<;qWfb1apl=()bG;lC?4(WMf)zUT7Bg2{@JHx=eLi@ zj?JgVSL^31tyf-k>{GewuJ)Lpbs>-^c>qtbbzc=v5l z{G$7b{$8kGyE@L*IonSZr?1Y#-uiT@FVtT(d-K^B#Vb;Y6h z#m3PWohv$T_nvMY{i5%IF6HZg?VmhxC?D;Uc;rX=f?uhb8_eJ`B@d$n`rFWt|tG`ITl+57MR ztKCcXYX5&Tkl*>E{lR6gf5olyVOK{#i=Au5z1r0)&4KFC`lahPS1GUekX~GM)Wa3G z>Mr}5YuP(LUHW}d_V)*k;;Y&FebQZD|LAwDy07&267_-Xs2>_@mw!>dxa!vBNB*9_ zcNO{FEB3Xo{3xH6=X;~w-@bd*-a~fhC?Cy<;@XeC*wshxF1^?|8m)g(yme?~zi59@ zzIXpCofq0S_5J;$#l6x!xM*E;Zs=Z@{Psbg`bPHB9L~@B>Y+OQUieGrXg&4251q%o zwjS+0c%^l|-Uspar8JlR%!j_yi}umJpgbC#bLqTlr}@jv-+0BR`F`)7bBm?)v&i^}qg(_4fkxf$XRs z8f%w-QNFn9*5yb3p1*e$`Q0n_wXggrpO)u)qut-Wd)3}UcIPM`&57dLkG|N|NAE7Z z*f<)ke^I=3Xk@==e^9=6|0|sr+Bfz6{iDUb(ml9nU36~fUYGp#L7)0Y_R<{A&-&`2 zI{jYwOXp}k^}7$9$Gx^5?LByUShRjkVL*@00HO z`bYKg^>?hl7pMAcXssqgO}E$)@>!A0w$b3^yK zIbJd1-d@mF6q;!)|{|_l3RMIq;Xx?<>u%zI^uHlWO;oz1qD) ze&>(&2baD66}QfZT^;=_cCHooYFDo`2dYQwm#*JjrM%uldU4fJ4_Dl(yX~0NM&zflb{hMA(p_KwsD9~ptorr)dx7GR9rZ(F?eZ_m7gyc7{K(()_pTzpd&R!? zl^^BP@_cW!``dS~+Iz_E9Oa`qQC$1c7rXlC-K7^BN2B#Gink7p>=*41%J=SnrSn4j zroO*_w76Hg2N$i2&JEq`lHWe)Q{Tv5n#1{7Up-W(-wS`~9IdB*_o4H+*Vd!G2d}iw z*ZUyezLe(DpZU<77k$tAM0qsY_tN>-PV<+SzwxyDu{68+O7oTaVYi=U_r2QrG>`5* zebwexUmm-Av*;f2SO4FCf0&Wq`J?^8Wv_q5t@B}5M?Z_5YsJ0V)ho?`>e2e8>o->^ zulJB%Ty@mL6}Rdx``ufNIE~TyOcWmv}E#>)p!%y?q zUcWf(`t0h^sNSOe>G^wCk>9;yU;E0B@@aX#H`@K}yI1W!WOt78v2%-SKl)-hyczFP)?H)bBoY9{1aNw0Gi_*73K~`c5y}NBe^EXmrk{^QxWZFE4-N zY3IB&yZK7~~a?;*Xo>Zpe+Zq;4(HP^Cte!BE~3;o{0SGzdu_epns z{p0I9?{{qbv_A6td$P26ai|}8xa|BWPQ69@)ARSPBENgZzV?+L<!Nc* z_qyb_5Bk(MvX|y?e%4nH)#>-bUphzYso#C*JnpykXy5-Ut@HIhh_^4Lx%6i~v_BV} zo9`Os(db-C`(8WEUta#k)8=2A-F&6_O8v0gU-zS$9r=;HbWUGsZuRA}yKjrV|JA)` z{LUZk57xdkJIY_@!!Axgi{@VbuXgoH{iC|He(CznRm$rbIY0KfTMS9xdMYM!UCt_o}^z?ADc!ol{)<(HFaS_M$#eooaULU$m||G_qf`KPbQN z|EqRhXfE}AceJ=)>0ZdEt&7eL^?~kz`)eQcDGu4o)rWfeQO7-G7gzQi%~`5z9%T2O z`~F{Po!5KNIQvqX%RW{+kBiRLexf`YyT7#c>5{(`=NwLHWFU^kfdtSBE`dRE;EAG{Ve?|3Q*>zv@ z#5sTL{bP4O_`9$4djHwQ`QB*t`|eeH580igd^DeP6{mmwAiH=ri z{XzNO{jYRhXdl$K&vfg{`zx)B&JC@P`f`u$L;d2|OMN>(>#Jvdb=3K~`W`RZM}49^8l6L_pW12ulHdJ#W%H`9I_56bVYk2T zM>YGR_<1svwG9#Fye; z+5GC)I^vu^_WrTEAN<`{dcFVb;(Twk`hEARy@%}1Q9kzk#OYr@$nKpj#r;ZiUbL<` zX=J}>e^7qk|5xq2&_1Ydzv)N2edy}`_k+vf1MAzIQ=Ykt`+xcSFhAR zsz>XWuHRgxyxv24Q65^ipK0tNcVTF zy8P9(^OudI7n@JpkFxim_mAIw;_trF-f0w1i}StF>U&qco7J^DNBL+@_d#6y(HFaS z_AAwqN256}$~PyC>=*4D%J2LCs+|{_OMUxHx30Xu(z@u}(E6w^_t-wvFOI#`xAU{U zdbB$3A-g!N-Fd2SJ=Bl-$nL(E*Y|I|p4Y3^UujPLnGb!B7ww}yQ67!%L218hr}<0% zuT)1~HM@IRn!A+8Zhsfu6I^k%JEyN)=g!`HQtcest9vi`-3PQkSo_lKD1V&~yEy$U zcCHooYFDq+KdMLTm#*JjrM%uldQl!)x8z6u_Hk)>oom?}N0-<8>b#na+^Z?yW}QSWAT?aom?n%jL4 z*M9WH?wu{g;qtSKLvvn~Z%!K7FWNVh-}nDjJ1;bs`u3S_U3q_{bFzjrO&) ze(f}W$^Vt=$g5^|Z%T8Q^4Q%2=U>f^{1@G;uXNt(%V&4r7M&x1b?+I!`+)WbYhRij z<*)N$7pI@a&b8uR?dp~KNA+m^()F9Gl-GMmFUmvfmi)-yJ}xbe1@ChwS38cIT~~a?;*V?53O7BBY*q2w7kx> z?2V(#eouX0>0hb7{I0Km6nD|z2VB1k*~P8jv5ga7HjYN|G_vEGulJz$kKcXb@4nLB zX%tV3^S#mP_uZ@Z9PJQbWoySGzYClmPjn*x#TRY8P@;k3rHm~}sWA0KNcIV#rj=m`WU(tD6S6z1R z@S=0%uXbON-+e&)gS9Wsj`Dk6wbS}p>|87E)vjKte^igwFI|7-(LE?W6YVcdU5z7rRciJoPR*f9(BZ zcR%>Muk?EV*~R(ZX!Z5)JgRGVj`FelB2NGML3ZzKDehOA^P+XlNhAA3`-bwp`(Nq2 z&_1YdzvHKS_`AhzLHWFU^kf*ZHuE)6Zh(T5+#-^-BGtdbED&`ps3!>pi3w z<)L*;e&lZ-mzLMLmc4Ow+3#t+!|Y$FzPzrle-yXq?@@XEE)*}0_B*zA{!-kc`1;v# z&DVR-`^WD-@poTo?=*_1#rfW7_51Esdk@*2qkJ^4`yj6U=!@Mu`<3d*qtTof<(rd6 z_KWrh<@f!6)y@merM`WpTUXv+Xs9NoG^f7xiT3BBbMyV8JR0q5sh`?u{*wPI)sa`t?%tN>F6FU1 z_r7;@DgL5!`%34lzI=A~ZP7XMSNERryANo8u=b_dQT{p~c5(Vy>|87E)vjKte^igw zFI~U6N_oA9^rAepZpn}Q?c>t&I@hu{jxPH>{Z;p!R*${w>mS9xQr-Sur1_D(?C)6q z_My&Ks)sAT_n`NW-+kimzS7=l6id2$f zoEPPrlScN7_6Ozn{eRWY3(cjzeWqJi-d|~5bZ%&U)R%i~ALy^!`zUr8} zREOREI`?Y!Me+ZN&ewYCvb%4K&XK>m_l)0tK>LHWFU^kfdtSBE`dRE;EAG{*-L#pKkKV!eRbT! z#))S~cJ-}?=0W`}ecAb~*YkSS`YX+49rK}eFZymSE?hgn%zAq z&0Wf4cMlfbM{Jz)sogn!s~oean=3jDt`9??GM(zG&{;)=ff^eKZ~7f#l70q zEA@}+(fXz9H&-dI_mEzcht@6mk-vRhT3+W`_Quhr?+AT2G_o&>^ZSez&))U*kKz}7 z-zZLf>z4H|&F*)s`mgM|^3+|F*L%?W$L~JzcVB7mG>WIi`QB*t?X!2Yx_0L%AI)t} zaqUN6?B3b0R7W0-=DaB1oHVjuv~MWC@BgcIUT7}$-8;H<<^7e`Mdyb4KVl(dxK|?BcL?=c&H+P+c^aKC9Wq)1~!#Uawkzr8%uqu;UW{7Um*bl$k`L;dy>?WcDc)uYAx z-e~u>?_Ra{klniSv2%-SKl)<#&R*09s#DEw{fpLBheq~`_6Ozn{eRWY3(cjz?~fMu zE8PqEv~|(Bp+3+(aDVNCKE)w>x%yB~KkB%L?BdFvqd7};&4cXTkFW2Abzbj5_EUjD|@&UtBe^Ofc+^}}v|OZTFDjdNaKX>RMtXYW0! zb|2ZRd;j=*59|-}Us@cBU+2RvPCtvCYsJ0V)hqRn>e2e8>o->^ulJB%l!w+W`K`;p z>ek-8&ZS=MOS7YQ7T51P_C@*F@6`KlUDQ9W`pfR`!{yhPxTV=q{lET>HUCBDjq5(t zZ$Hs~dY4f>TD4XkB$^WWQ*CP=4S4 zSM9veTwm%o0TfZB5wYvwUbEuu>FE4-NY3IB&yZK7Ab$u-0I6`?>(t@AK9yW|M+_k><{u^S{#aB=ff^eKZ~7f#l70qEA@}+(fXz9 zH&-dI_mEzcht@6mt;@ga*516%rC#kzv!izwy{Faei{i2GIc;6kKd$=AzWyGtqrR3_ zUmS|ZUw_A{hxVg%-ni~V{q_^>r*|3Eqs9B)X!o}7UbXj--MaFzbBb#}`eOIaUepJw zQ_XJui`G?#M)r&L2j%zuf7Q+l&85D-f3&z?>0ZdEt&7eL^?~kz`)eQcDGu4o)rWfe zQO7-G7gzQi%~`5z9%T2OdsomqLYLNgy$6l6FQvKcWA*hN@Y_dyB0r7pLFpW7r}<0% zuT)1~HM@IRn!A+8Zhzf_YIfwm=)Att{Z?N-d+$lLb7ZgX{o{8Z(EecUOS7Z=bw2Fk z^t0HxR@|#yy;A?E9<5)xesh)bdJpMEd1&2|ANkwIrR8<5Wp5l^dS@?sXT?=}zf0dK zd)L=Ls=Mg@FR%F3w>t90A%D4k$EwqK_0WEly$8L2{O%Kf_m%ceqj*}J?~PW!?_Ra{ zkli`TM{~Ll;@XeC*uAq~sg67v&3RG2Ica3SXn#H_r34HrLEWVde!Zf*^zvOqnUfI0rtB$!#b=d81?*n~N{=cI0wXVAC?%Ses z{$TA(v!nc;SM9WZ7CYC9d$p@q>L1mk^-I^^d1-mQhxDR6v~J0d{50wZYnRu#mc4Ow z>G#w{zstl`d#^8cef^`luk`moJYDLWz45d-6u;QtvFbEly^GEtd;i$o5B}~ez21Lz zalSWN{l0tE-a~fhC?9(-#OYr@$nKpj#r;ZiUbL<`X=J}>e^7qk|5xq2&_1YdzvXWuD|or@_G;HMR{o5k{|hL z)DPA!uX8PXLDyMZaTPSDtzooj>;e zvAZAq-B)_O|Lo#?Z?yV-_o}^z?9Nd>_Fjn7zkZP2J6nqTmFB!?U31dNe$oD*{J#IM z+IgXUP~U#jtt;=Zv@SX~w7&hJ-DCT}u5V;7_3iwuub%bQaSt0Oo*miMw;pyL{jpou zdw|{py0l);>s9NoG?)I&hxX^9bMyV8JQ|%tsh`?u{_^rSo;K&w?B*-YSL%n|Jt*CW z@-@!6eWkgrBcHwZq}qLCukO9%cm8O9u=b_dQT{p~c5(Vy>|87E)vjKte^igwFI~U6 zN_oA9^rAepZpm+5{#Cd3=5;RhYG0Zi`+oL)r7yNlzf-^4tc&``)~Qy9{pyQf{8!2^ zSDpTj6{qiu&KuW#sNa5~{q!!QdbD`o8|~iq-K+Kb*kB| zf6=M;-T& zU0m67G-s)t*mds6Kj*{gg1_}vGzKUn+H>?ptIRXeSp#m=?jUhV3Y z`bYI>{nGV!URqx7A-yOMty}UVKaKjq+U0ewWp5l^`i`*gF#VP4q5Q6|e-yX4{vP!A zWwku^rP)#ZE9LneD{swHyYt81KX&(nzxztB_n%#y?~PW!?_Ra{l-)VXM{~Ll;`FZ{ zWcSXN;(nz$FIv}}G_qf`KPbQN|EqRhXdl$K&vfg``zx)B&JC@P`f`u$L;d2|OMN>( z>#Jvdb=3nLZ z`Ahzoom?}N0;7N^gH7#<(Ixw z_O7pg?0dfS^>;v?IP)wmUL1;lr8?sMj#UrsN7;MO`^WD-@poTo?=*_1#rfW7^}Va! z&Fb2nqkJ@{`yj6U=!@Mu`<3d*qtTof<(rd6_KWrn<@f!6)y@merM`WpTUXv+XhGu zk>7nl`-8PF&5rVWUbWNuS?pXZ?$xedsee?D)-PRu=cVQK9@2~Q(7Gi*^3$jvtX*E` zTK2}#rF*@&zO%-QTl9W+ef^{F>y`Rg>+)M)+|umWKGge4zhl*1bpF`;$L@accVFrC z{_$!?i}T#eMfQn*AKFLXG?Lv(wrBqYfc*3FWNVh@7@1O=Y{q`efvze zzP!KEy6D`{`u2x*kL?4yzLCAuxAU{Ude&FRJ#3tKc4SxIdT1Wh-_qEcJ)gAqk6P{>H5u8%IiI(7v-UKOMc{U zAD5Qbxt6_gbm==n_nb!dMREOIR(sdiKd$=hzCUEQPRXw?{-xQ`I;h@Ps(1a4?LFxI z<9DCEcJ)gAqk6P{>H5u8%IiI( z7v-UKOMc{UAD5Qbxt6_gbm==n-w}=Mi{kuF!`i#P{&Cf3U*D~DP(12qX?fyst;eq3 zqB!+%{f_ND=>6k&pZL44w09cC)8c$@wEE^UM|JJaQ9hd6K8tHV`eOIaex*9{Xf)?V z`R1gN{i1zC`F;Oiwev!Asc)a@)|K~HS{I!gS|9c09@~fd#j%(Ac7E1Zk5ei?XUB%W=H;u&gm=HxwH44R69rZ>fTF!_W|t>*1j}5%3tThE>1s-oomIt+SM!d zkLuC-rRz6WDX;gCUX+K{E%}kZeOy{z=UVp0(WUPQeMdC1FN*U!1#9p6`o~qDeSN>y zLGh@crR9mkwH~{Ai{jKP`#ZMxp!bj8ed6!F(%xwlPmA-t(dwJa9M!cuNBP*f#kC)O zv3qB~QXP3Tn)9N3bJECu(Y~SlzW=YbQsO;;?q-slN5F{aA;cRv*=+OY8N#UbX&8>*!A(Xn!s`H{UPHqtU*W`l+4f zFZtcCU#X6Ht~%n%*0H~f?gc-M&Z89nmFlQ3pWS_1bdLPhy=VOH19snPUz#1&>3P*o z>u0fZt+-dadZqqRJzBqX{hgPV*Lz4W%0ugx{K!wEez10Voom?}N0;t1y4N(aFN*V> zpx-TY*VjL;`s{tr*wsh&I|2>`tBLsy7K-?>!Nc* z`(j?&J!aRZzLCAuxAU{Ude&FRJ#3tK_ELT8p*%EKb?fohPFt_%^{Vw(+K;uq{x3QQ zbxfSSc>o4lV zIUqYW|J5!|KVRv+D6cy9sa$o~-3!_}wDT*?S^K5!ciH{+T_8V=-q}U-mF<(Co$fxc z+oy}pLtM4`sJ=eful-hcvEK`=FKmA8`Vq&^POHBte^H*gi}uU!*mVz<-M*rGMwiY7 zt*<}#T3+kXC{F&W%a7JU^?Vn#)1^G*r`1FDi=9JVv~R!C{rJk>xAxh7_%F(9obzR; z?Wg@HTaT94dUR<$cJs4$UG>*`>NKzSLH$y``pE8n_jjjruAR1yxUaM?s1I|gU)_CJ z+P=z1@${ngapl!7ZrS-S>R%j{p7d) zqCT7hvZMU^*|G7o`(Ntw)h#<+^t?pvK7g}G~{Mz*+j-Q=Ye^LITJarfCm-GJmJJ!Bp zziW1ME@*xIxu@b=k4ACwS6zOz4yxz7sGTn5AwR7ivR~{R>Y{!7mF~w^_P({x_QU^5 z=Xue7+K*Bn)}`HZcI(h({pNRmT~}Y$qg%Im>Xp_O9ZFK*fSFX~?$jjw&T58CZ$3B&-4!iG^whryQN^{nJY5QGvzkL_TPowu7Yd2rn zKKa?{?nCYTi`~y^_0anEm!0M>)m`kppxbxz`PubD^Rv_HFUnt(r|zPCaQ?V{$M$~M zS9H(l(t4;bbGz5-EH+Ml>#C39mRt7m=ZtUf=D;`wP*e^DJ2XCHf>`mUXpPk*KQ*ze3f znvX7*zjpny)2&DIqdb1BU)-|uU+g*4X#LW_x*p>?iH&0 zulTw@_SJrr`caQ|4mAIw`^Rr?_V&^G{MJML^3$lkbxU=Sf1R(s*E!2y{>!dD>Z9{4 z&A#@dcKgFlt79ItPbi)r>leqa{#WW#9*x#t)CWI}_FJFGUh@Byt*gJ*Wf$+BEc&~~ zU)}r2Z+{o+zef_1Mum@=MQmp1XTPWpig!Q0 zviD8hMdx$Txm|Jg!M@U^e$=JiXPSReKm6vmUi)y)>UTZ%mCwGY4|R~;`S!mb;GAox zttaj)?F;HlUHe#Vzb@KG`;GEwv>w_YY#)u|M|tXhW&5Pj`iuJDr_p{dx)<0u^DR4n z*?9elT>(GY74nx%t`Kw|e^LXQxYk z9T%vcRgBP)^nd* zS3H^z*~`|kjy`GM|0}KY^*)HVFQ`v-^=Cfhzi40Op*$MJqd0Uf{8+y@_Qq9TapGU8 zUfI6Z{<4?$>0;NdZoaw1@mG777u{q2YWEEJt%ugdWv_q5nGfYv>*p(dH|15wK9#Eu zyL&-fhqe!;xof|)eL{6ge&na!|3&Y7=c5-$|)2 zY<}(f5y#I?cb#hai{jN?v>(oUvG>I9Sovr_-M5loKjx(y@4m9r$Sz-f{jisvM?LfN zv(qI%^3&>}{$6Pgb&>s7x*uQJ`(}>SNB8~Gt4{6q(|(lV+n+dg_lGX)H@`UbP+!(} z55%K*6o=x}!OL!c{jNmkOqc45`%3$QzH9T?$EDd9o8LHo`)MD{MI(Dz|Jv`y(WvgF z#aX8`S1FF&dG_4c%f>D392a}u<`rK*ySYpGwVRh+9J0F??g7?Lm*RcLw0$m{Uw^eZ zE_>_BTQq+u-g?#Sxcuz6;=6uz=d!=-)z)EGPahZMU9`S)r?L63|97$f=f59H;;Z}b z7ksC_M}F(C{?uLkfB*G?;?zfdV(t2?pIyBAi}Dxcsk>;Ou6y7fp?lf>g5puU zb8%15JBQvu^=j9be6*j)ZeIN$d)KiZJF;W#wEFr){gmo;ZuR(4Ts1q&|CQz_-6OP~ z^DUh}E`RNG*SDX}&wbYqEzbP@Rz@b=cL@$3=M;t?%4vZ2tP)Rl9eWz1n+?zEj^NyY*Lp>Mr)Z7pFez6KmIB z{p{k^UzEQnPu)fPbln5@2;C$5>36L2_O76Kl;>RJp?40wgX-0;FZpOck=?xdLH4d= zJ$7Wr+G+LmiTWwk>)h(`qqu5zl>aNuQMyNHJ?C3Ge_a0B>8@`-ouB)zA6lIG^@-wX zWVfz9)Iomn&Xpg<(asIkNBg9%d63;bU+-1zwEFt_O7oz(&0`;{%~$FN`Dt8s>!9_~ zeEe8{c70f1f4F`Jvim)V?D{}ew%TcQ5K!|LP(8MRooaonN`?yBDwTQ{(8; zzS@uK)t`H(-^NwDU+mV?PuF$M>UTYMb(+tvzI#IJ1J%=>tgNdbl+&SpC!NjKz8d{r+V4NE$YL1UumCPpB?$#Q)Dk!AMN+j?t{3lQ{8^6 z?KAS5=ausGCyyQZUny_VyvzU9u3l*#RF^gnUBB~pZdiM@@1@kAb@0``;?zO&;Ii{$ zs(0fE9`=U5=p?<3M;OE;ir3=@+f* zcWm!3t)4tuAL?HekBwI!#rfVA#g+P86wgni{khn8qI0djYnR8a?z%UZ-9Do8v`*Q2 z^4Mu)zu0wXbG9DsT=eZ+O7+-jb&wq!-}S63p1)N8+K1)$PNDlomw#PP{#UO4YggAC zSi3y?yy!b9#V@<{ex-A6zV+PSuUvK6(Yo#_vX`rmtKa(FclNGR-F~Z`GxD3~mGbl_ zj~)46DR0rd%m3A`UTGdwmo^Vwzw>u)SbMcN^j+!CI{0c|aq6IXaM}5>aZ9^*=slv5 zeNmjbP(RiBac;Qkv-iDcS0D9*wOhY_cJb;j%3qYH?xKCV*nIn@U$n0Kx_-y1Cy&;L z`WMAxH=3MvYvfD>=&RvIYJ$dXjvR_>L zz}|YabJ4eREcMS$tAo}9-2s-6Dpdh*N8Q@#4HUEO{+waat9 z_M1lblHd6tyY;M7z3k!^^*~|9fezU)7aptt{V!w-O^C16MT3>(i z*^&R1>MW|e{9om1eQuYGBD>^tNgU3UIb ze7`&Muv-`PkF8tX|D2Xzebf)uZvX3N7q9-J{6%@{F4`yivM3Ji8;$zuz3uN<>&T;j zrG9)zXg_hqiAQ~*zG&>eiesnI{#?|DeoKDjuVzQ@kp8~X+~}Omht_L-dF(W@qjM|O zFU>FBKGW8*4|Hh`_GijtlhrW&n{m5Mfr>J)Lpb6%g?SaTzS@A_t)>( z=2ttHi}J+zu28?&xTV!C?Z={cej42$v|pw7!W{ZUeiT>Dj?Tq?e5JWiJ$<3|T3;SJ zjqJGYw{^;{D_=d-w{t1gWvAOGJ2t-SSy#MzrM^)e`l9=V?iXFE|0~-+ebuuM7rU>G z=Xd_rzbFsI>zn3hxBu14-Z&c9JhhuoKm7hVDdl?~X%x?o?BdG$=|%bV*KU2}Uo?Ly-g?#Sxcuz6;=6ved|Y+-<%@HFocl%Z zZE0WGmu5%t{&z>{e`i1=`=U7CCHg*T>zDRr)o1s=(?xdmQ9oF_ccOlF@#-(iUzDfr zqJ6sVfqR7Rk@H0NuJz=h-?8$Xi+t~0>AhsvzjIUv?I*J9&m73!b*#sZ>{vUkzWLF- zrFi)zKk`?zqxpWNzDx6?_3StHylE7N&b2frwjby1-qYgD@BWnLV5jBl$DF7i@y^wH zC{CVuR3Gg}=P?(%_rN=a-UGT+Uw>a|Ur;^sbRTK)7xg2~zS78Ue=n+w;++f4&u*^j zWp5mf`Y+AHPV0wu4=%bN?hCFw^OX8#zqIofr++lxr5o3H`%%jCJ<%whAKAr~_0#Tu zsqSU>9-=(;+Xv0RxZ>0=)scs_)6Qd2KFTZk)!}!4e3uvPvwft$^8d4TFUgkVxD`b; z0Db2_S#@G{w%IYzZ$xVw4nPn9>3laUt0Z%g$9sz2Q#R5U#rZDLcgeQ@;>u6|-#ZgmBf6@9y>*QTDr|TZLN9Z2)9=K=CXC1l+{T=Jxdhbf_C0+gNT+I{J zt3G`oz5CdYj&!WfmS2BpUZr^JOMc|9O-KFxO6OPVAMIz}=-g{p=U?ZlPi#HT+r4Lt z)4%&u>VwX9PU_Mpsz5a)rkk5Ou==vi(>*{s($$Id(mgHBhkn`m6sJF2ebwK% zQ5?HytB)7kPbTEA$Wyo=^^ z-2?Zi>^$8=_YdjVd*gR3zk6D$ldgXGO7lebPhIVY-u&#+{;0m%boo)grMj#u`H{ai z9qs!oTPM47v0v|p^P#hmj^4*opV)q!ciDL7RO*M$ZXNn8#kVfDI>gxryANAG{qC!K zr_g)AF3sC{ex>@*ccvfjPi^~_@+_U-yf)?hG^f&fw4QyP zm;1SBJwJQV+{_Cvo!>ra-izWEty8Z!Ha5 zuD<%|;#VDX>lekzyJ${pE_C+?oijhuoqOrKS#f^HibLO>eAU^_Yo2I*>$Q$9PF|!V z9jmkDFPC4x&PN>5zf#_HA1=N7>ASi-?f0_He{rpMFRin`y4d2>Z5}9|jdYa9d7%2l zd7tfv;;f^iKE$!*)eq9$^Yva;XUi|&SK1%jmpST>{YvZPN9);WfAgz->E?mf|1IsS zpDUlZvU$`=M>^8oU!<4(=svT{<)^E|JzaGF_-ps?9>4p5om=%~(@~zDS9P{}7Tec~ zd)4JDeGe!fTYv2O^|NRm@~{`1$GURGRqsCh^wk%=Z2!xaAG>dL{!-lf?&z)anpgK( z@$#d3u)6QJe!BSXU)%ac@$xQuf0{=f$S*Iwb7k|R^=$XT_r{jr?^y5NMf>)g?1$pm zoeSMu%Eo`S?SHX(|5Y9O>-+y|J1=yut3Gk)oXgAK_|8XPrTXY>>wB;Cv2=AhKY6UP z58HjFH=bR}-@esO||C>ed#?9>;IPa)z6hr zT-iM8r6V2b?kCbq{zZAwchh*fxZac6?mfMB?;pSWfSp_QWz$ifo>z6YdKTN)ihI@N zD}4_rA6tLy`t`GD9`djko5#9x#nD%Ox_i9(qL=M|+3L8cFLAZii|Y3Muv@1*uUzZx zi|WMc>ZzYD{;JFRMeF5V^qwp~U0rBh=fbuQtz)|vzBjghymzm^W9_SN?7rgEp%0|H zKd2tI__F@m_P;3as)v429u(jA|J8P0Xuk57;?Oylm%s7dM_;A-=xpnIuFebVSEuun zM?4$dXZzDvJiq+yTm2WUN8bZ`@%0`w&Rj};qyE_Hx@eBA+q&v&KK8x*;&8=()#cNV z_3ToArFHampXg}*?lsa&>rg!VmDb7Mds5rIr`PWNI3CttCwBBzDnzQ57~<=k9@e|s&^lL`s!=x?VnxBf3fw{_8y?`qu)#SS3kO^^4ZYyXgIC9(5qUy!16!ezcx#o$rlpzRnN*j#V!moul<= zKmChqJ>uy0DdqW<_P_XR{%f5+`~JV$&I{MP)FJMb?!iTUpgPe0{OtPOpNsB+KAd0k zv&Hv&QO8C3I|unMy8oySc4?p2d(b#@DZkd)y4dPB7vyK7ea)GT)oVLHw9ftcmGbCE z9(u`Ns+VpZ8|g^59_gk27p+I%QRC_2dQWP*kM!ETfBfc(?Wg*(=_rpmVs*B9zS8|I zuRP{dt~~U9Klb6*KkBc%{Nlxx>ZrbK`4*jzy!v4;Hg4IiOTMM^m*U%p?}WbASzmu` zeIUR5C?8fg*ZS$=k-&1w1R>O$-6gXUu$TE`aG??s&V4!wtd$5xlmJk%o} zic`P5xb*HPu3UM<+5e(=^RhqEFUo`B`~JV$eILzLezcGC#nt!a?>civ`=YuQ#aWM+ zu7CH-I_F|Nx{u~TU-A6%w{QK(i`Bbs2{RoO`n9{_xlC{o{B3Xg;{~`d8dK zAG$p1S!`b`?p2qs^nIgzY<;rp*UzH9+Or`#Oms)pDzBY%lbv@OQI;?$uJq_2J5t6!;^osq>;dD8BFitL?ndeC02l>nokl#qOi8S6bh>)^-o* z`gVTJ!$$K_7hS$Z`8x-B_hyx<1+U>qmcBy|y@XU)g9~$&dVOlozY7zUXE9U$*@&ww~JF1N7bW zdug4{4dq=lKNM$QY(Lf2Cyt-amVeRuMeF2U^xia&I*?yp`kJfwMfbw@#x`H)SE`rp zcdR;`C#qBb;?%{Dbo-R@{7Q9QRL|On|6~uT7Eeby ziog7gSBH8_@64+%Zmr}0y03f}o8LaAIP+L^U(kG)-FxF6ySIzhsjFP|)6K;?I(qle z{~Z@L(ig?`yVDolJndhu{B-~K6Ok@Ist2o^fBkgv@-JGyXq~)^=H$Fz>AsmO(tZEP zkDb4J;y%CDp?<$(<#GSa6Y2I*FB|1W^{}yX6-Q?;+P_psbv8d%XP4_doXf?YoA~x; zU;Cl>+SU2>SIUp#(Ruqm*y?X`=dCd+YjmTyDum|TBo0j-YN7RuuJ)S zzRPxA{N{pm{g{7k{arM#700g+^<7jaif(yDVJam5#+4f=AzjSfdFFLPR z_I|my?$e@m>T2CfZydYyo}u@cjr2uv{qFQh@BZq?m7nhYN4or|9<1&ish=)h{zdB- zt&?}roSgS7-M`*b>-!xw&Uw(;)?N3KU;Wnm9qU}p6UEs_A4qS0c4>dqe{Fi#(@XVK zXY*rqcBwz>zS4SheycBWD86=e{@41%qw_A!L!5Ie^+RV{r!M(Q@#3sQx;XnFKU$}s ze)rYAQ|LWlm-5^1E3HT8sUPo8?apo4=5+bht&YVjjvw`-5A&|Qba9LFqq@CUbT&U- zUi%{*>COS^<(jv-@L#s`6Q}O7eevs?Ee@A%zsAwgdK7>88{ax#ZE?R+9i@E!9@eHK zKQ5gg#dZJM?hBglvU_jbWA}E^I`yGC*+^F}o1f0^`|EvOI)5p?-<`hb_D6MK^VaTr z&o4i!2dkTZ{dDp2FIvB7oxF?Y0YYo5q&AN8_P zUQ`boJ6CaZ_M-hubyR2bV|8}vJbtC~D%+oZyN`8rHqudlrTi!!owx6Utxj>yu@pyV zt5ZGpM{!8EAJXM_Ur>IuPCpmddsUq+|F1NcQa|Qjn|{%{_HA8tb1j`qbvA!_`5Vu6 z&da9juhd_u9=iARqW7=o#O`|c?<@7~K3Pw9?-u?2;;-F*f56{+U_QvtF8Q(kbv|_Q z>RD`GEACa7uha+1$5t=9etnhJ^&YYpS04Fr#jU(cUwtjT{jsea)30t!G>3dti6!|uI?PI zN55nBE3WmZi!MKUciD@LW25~qink9N=@-ont?&E)YCA79H~HN|wzyZi2N&&&&JEpb zbpHJ2piXroz0`;Ev%h>OPrn!b(mC2s{?0)^s4x4mf63o-yKMc7t4-HOsh?6kboZfjALzB63xBCzbZ+d6;!(f7C$-&2 zdhOmne%}q64=%m_6}QfZE{}Q^+t-SF)#WSof%37{%dTHvrFFfB?8TKwJ`^Vpz2D`s zty_I9z5TPxekaTB_f^~XSc;=}fAypIMc-f9`1hjRV$yb^S-Faj0 z1s(Z2*LBa!Nu3vypyL9cX>u z|5w|2p}EQLoY>A4t*c#~-+9%xFFH4L?(R36-yGCis-MnQr@HJfA6n-fUUAM(ew4q| z59+J7ystFh*L%=7b1C&>9<|Nu;+mgz{IC65x8nIx-OkhZTYKq^W25}2t~ED)@%HU~ zp(DS2kzSf>X)g53?sundTz%Hxxcb>>PG~(F#q+c4U%EK!`#rMD_5ISZ>+ExJt*hRB zYQOGdK8x-Nn%lD7FLCax`*iW@lV893WFuX?_Y%E(Y@{!W>vyM5y8Rbde)@W^|d&Hy%+V1qucLS zns3*)PPTo^)@2_0M|#PBQN4@pL!ILJYx_PH)yrSocZ2-)LG@tuWz*4m=YrMQ>iNob z9#6HdZWj(yR2nFkx2kF8!b zFSIWk>H4PIzZ7quQoQ?Azx?u)`ayp6d8g2Oz%K1~o&WON7uBUN`D?5DqV@8kb!;?W zG-q^A__2O*banj7_Aih8rT9|4baS|9KK8}6uJzQmKdMU{z4p3CboXf4?lX$l585BA zFPo0m>jSH^)$^6kr@Zo*Q@Qfcojcn;Z2hDD%F8caT&a%g%a(7^`N*pu_G06fZBA(J zY@{!Wvk&T<-Fo?P<)`~jkS;%}2dk^2e!6)17p-5kPTocRJMYEb6Yo^>s{hI(PW^2A zbUod^{B-N}qb_>s@6!4mD?fH$y8TgHZF=iiw)Of^FVcUdd;FE|6}C?Gn7{MI)$ir+ zI&(ntMe*ujt9McT_0#n&ul>>f;<5R}uX`lFy7AJ@vFEDp>TG%0UnxJTL*JbP+rAgY zn*&^$Ug4@&W+I_c(c(R}QSYhCN9ZGTjkIC^dOV^N*_wcRu1 z*ALnst1p|5*6RbSv(>ZMzE<3;E??Pmm5;7JcK!On)wewSi{eWASHEm^p?J2u?TcPE zZrPo)d$@G|QoMaoAMDo2kIh@V?>WExs2;4Yj{521k-_3yl~_m3`~zj@Wq z-}j6zPW^2Ah(mhw(bb_olo!SMyM%tnHoyIsZhsV4o34(F*6Tx^NdHQC-H)$yUT8k8 z!<@z8>ihC{oq3uE8^x=Gt$%g1Q6Frit5aV4qrBpgE??v6D8IV#($%-#tLkidzupJ& z_AT{e4z=kQty3pj$42)6%^%I5AL|!KSBHITFJ0WC{Y(3nt*iIL9Oxzg#Z@m|zo^do zYxkbiRxiD_dx!k?LG@tuWz*66bv|@)>iNob9#Z_kFUj9Yv7p;?b(fK&< z#oiO|XXh$k$o+wWJsBgbxyAQjxAIejkZvTtz%X+-lFTcF& z-dwslqCUHyb1fUk7RSEm+!mXU?flej9%x@S()CSu9;JBul;T^DIQiu%^@IHCTklnM zw*9X2Uw-?dy4p`|`(Ct8pJ*K$>DIBadTr-|*0F!3bGM&7@|XH4t*5i)Lwd>ow`^Yh zH7{Mfd$;KC7k};k_rLkgW6`}_bLO{RTo*V7UwQe( zJAX9Ck{|im*n3vJ{jpzZF5=nxx!AeX))%&Zx^?PcztXz0^*0|sz4h?Z+1;nMe2e1c zUF^Kv2k#Bi#nZdL`nz9s^|=pdo&1Y_xA0r99{J@%ar#;GJJvq`iso?9oYk#Qq+gT= z#rOSxwfjEmL;lh^ywdqyGzaXyZ0CE~;!64H`qwvqX+Jb~wtVYe$=^9xClAu?cjU#H(&L#(LRgf%2lsC{A~5G`RVNLQ``DQ z@$xR3S80C8FAv>5s2*{sFSc`VA5fh4ZqeU;e)-ga_Cs;{(FfAy|N1*tJ=U|)d@ia( z-PR-hqC6Aaa>5X?^Q2x%r zI(d+8pUzvJ>TLekd(b#@LEp1^EV@4z^<^$-9UD6rHkuc^#e4U^(s`kL@|Vu_mCom)`C|8NUA4tq@7&A!_3ixlOZn}K z^v2hgzjLr&KBT+vowq#I+5BJcgLrc(uX(EfV*Aslb!?>|^)pp+KokhC)LANj35ADN7b-E8I-hTF}UEO+h$S)s?lW)=Q#InC* z>sJ@Ow*D`wSDZXZzbFri@B9C1_kEO4{?a+T()nF92j}Nrmi4pM$?iIK{rcB;>u~P! z%OjpHulQ1Y=U|<@NSEI|f2Dn1??L0t1?_7-s86LWp}RZSL#!I zZMwQj{gmpZv*knjMe*geuJu~qb#(FW$)ev;{Iz@k_|0R{Jy!>c=V#Z?#ucZYuk_uN zR~~aJS01`|nQb4od7ysF%P+q7!W_#LSKYlres=p~ztSAUv-NYabE*B>w{_}ZztXyL ztrw?1am%Ko^{lXc; zMf=sH5d1?te>sU#@7~4*T247 zhx^g|;_32=FU5Bb)~N^S^1J7+w9o53Xq>sAef4Edwe@*XU+P5b*hsgIjn!+LD_UpH zuhh5r+I0C#{gmpZ%d@DD#Z}jer>g_)&yRF*y+5_xBYJKB|9Oz#JQlq}YtH=Ei>tkK zaq9U>zaPshk2#eq58Zdlwh!ApP=Dp+7w`Ph97}%Wcb{I_{@AZH2k~rmq55#?{3s6F z7u|f-!AA346o)IWx;p90rla`M{-t$GH;2W}&3l9V*3-M6dewu~orC*;;^klLduzRV z(Uv}5M+F$R3cymGDnZDG=MsvAnZsvm4v5{^a z8>`ngN3_n|f2BFdE06r8eoFPyoO`n9cNBkZ_YV2ZW6^i8 z=FD%sxY|n>r=G9$-IP}zb1GLJx_6mvAGUd*{>sZQ-ua_Bmi)-?KE1O2v0rHp;@Rp# z_2JU_Q5?1}y7{Vujpn;34p&@tb<&qjNAacoOY4?y4vU?e_Xhc`r*}W~st2n(2loNR z%fHz7)_V2GuU_?wV=wyqiuPlde#bVBjplGs-D@87i}Ij&@BUZzebk5irE_?t^Sju4 z(S7UhIg77eTOPXp_05mYS$;HkwtT%uZ23C}`=C7bWp~}H{q;VGHy8AssSEXo=5bM9 z=7`p@QGDmL^v0zvbl@Z%#-r`H`QE?iE&Vf9%CI7xDGe z7aL#O+^~Jpo1c#IE;`q;dApvU-u?LL?Cw)rzD4o!E}EC~NAD2Q-4D8Y(LVBuXQR5+ z$>z78eQH;?-u|tJF3$Qze_#1))64Zc)_QeZ?A)wJbNZFu{jco%cmsH zR9EN7Pp?hK`t|Sp>|d_^g_)&yRF*{yAsSz2dL!pPP{1 zJQlq}YtH=Ei>tkKaq9WXen;#pk2#grqdaWiDZ762K>d}MU%c~2b1eCh-+g*z`(wY- zT*R~0h3dnl^P@OyUv%?T2OG_IQ5>$g>guE~n~vg3`H}4JdTTkzP>QxU` zcMk3YikE+}@2&Ofkzc*)7sp=g?-Bb|uPv^>W9jB_QQd1E^o#PKc<=sK_I=cc{H1ev zrSrSkd(nOC?>URFURxfz{`Jj|&RKpmceZ@JKWzCs2m7Er_GNe7tNryph&LDXo#{t? zY&4IH`Z7nfj*a3wr=>Sey=dLvvin^5^jFGLz4_{AH*VSXS?oIV63<`T{aEx~@YnV` z3;CTt>KB(@|BBNOT31^=U%9^1D~~yq%_k4N?@ayXfch;jzj$*(ddZLcZ1-u=J#2sM z#Wff4_0ty{U)$WUebSqsj`A)#*RpxLo}b?R`04EKQ(L}8@$xR3m-AliJ@NNPy=WhK z-8VL>Tb*ov`>CgPb?fcldg$V;U(_FeZF<=_zhmj@xY)T_kLL6%z58F;_wm}_ykF_u zE}8?nAE>U*k)K|hj`i!``Psi*`^(#VQ(OMd!Fu_TZk;;Ocg8O5^YuQ6Z$A5(-=exN zTCZQUj*WEd*jRnp&ZXl^p0-Tl?oe01@>FSXqtdTsX(`Q=~q zeY8ID{A_Wxmo82{U%B4ND~~yq%_k4tch9chJWzk-{ptL z_}cWv#@9ACG$;Mi(LQXHchR|)@?qCk?|%GrcK4~>eQJxBchS6@Klbl2-TmNiUir{D zuu-3_TfBW1TaP^Ik)Ms?tY6e${d8P@x;VdM_5VtB%ZK!f-ABB4|10}GHh*pB@JjVv zG*|QJJyD1I#1_xyUzC?$|JJt-Hox_3REKyrit8NI!H;zD>O}7WyR`1>eGuP#{m#rm zf3-V@+SbvrI@|uJe{>)DvHr%HbLY;+D~=!Sv*>*4r>lFh=hL|MU;nb3cjciszV^CL z-M{PD=8xvfMmm}!(viL>uk%@P%fEDa_&bN%=8XKyrlWn3zPR$!OW*xhItT3f>TL6S z<#mqy?kW2BPPlY_6o>AyxzW{cf9He#J=o%Z|M>xe;?#k3tgfE=>5UhM-Jjld;^bX4 zFXvsF7u~(7&hPJv`Iq*y9-WUo>wBR)H#X8y-}a?1sskHG$K|Js)2H9D=Cr6@^~;C! zi}K(VfBEHkrTc=;w{~@Yzej4Dt9jI3_s%}*Z(MEniq2L~_jS(lcR%{7n=XI*R|m?c zF1GLgmFkrDqJ6~KAI-sh==wwb`3{y{{ff6AsuS(ckM%dsob*?F>5XHfeHP8Je!99B zofkH){nx*2`Ecc-i(B;m)ZhKBufYMV3iFPo0` zLHgp#PcOY|U+EmM>#MWP@0Hg%^1G+#_ZKdmAH|`2Y;JV*+u!-1-+_z$J0MOSNXP2x zsh{3>ai|V9cAYqR7tPCgm*z!xZ>sbAyJG&O{j5jlBhUI?=+2Febkw(f>5J;X#?f*4 z>EiThKG*M9^{QV!q+gT=ulUO^&nw**biTE#^ZVUW+g#0~_PTfWQGerVyH|9!db+Q3 zmcRSaSKV~^+rK(cK6SBu|F2Z1ycg{w&i-f)=0n#X>aX;jES=w+*y8Pn>O}kVWBrXY zC;io4dgIt=pG9-5pRVpj=Y@@H|Mf3hK3sX|;;wx4TaWd#%^%H|UGkeFKOOlOt#dx? z()w3joOSwNv_FblHXZGQba}8k+c_+nXX!g$di`v)uD1ETQXc(x9e@8$`u&Cc>_u_d z?}M%m`>PNA4qWWt0dcL1j{I!(u=(k1{!$&)yH1?Ei{@o6W#{I8(Antk$hyDuwU4~M zlSSVPKihetzV%5LU-ILxe(U7Xm-*oO9jjjT%ZK!f^8710zjEbwFZ`}xFE)-{nyY!# zUiGTO(Y_ls^n^><(AEPwZ-%hUCA{kkV?b)bCeWc&VKsULYR+FzWyQQ!Joar$bX zY;4@p#n}(li|XXZ`WvS&=UjW~jbo#I7R|AKy1EzL3v68buYcL{;mSi7cjc?!daR$_ z`I|2t=}1Sq@4}oHt>b5x%fEE%`Srip`yq~An~wHDx^-Bc?Hm@(v-I6ouibOuug*5V zS6Zk4uH*0DUBADOpS>s!`+d;WVSn{ue-GB~-vR#CMMr+Ndf5DQHh-y(>Rl&J-bM2= zm$GwnU+8S~cV*pQ`n8YmWYPD+&vu@uZ++6mm;Csv-#U5B%Y4fIj#aPvKdHxlh zU%B$T7q9PA{fQ#i<+ht-lqgud;pf(^1`Qv>&P$)ya?bH%?#9x%Sc<$42`snj?y1t9#MCz{a)z z`j;&qt~_*cSHAkK$NJfwzxgh`esS&vRb90aAZ1neI-Cz2(kMCsB_rlM1o~Unq(#4nj_^aPKdCbdv z7X6M@ulnUf`bByE6`fzX@_U!iJ*i!t-<)cjt9jI3^|^QIZCq{li*7&lbYJHzfA^!y z)Ae-u+rK(cK6SEv|F1NEc`n*toVrop`de}GT(n+)XdN5%qYgIG%jKu*?^mvQsKfpF zmGWcr(Cvrx(mgEQ%d&B7^Aqnp7tNQywsS}RuasZk*3*&ymGUggyZm2u`O4NSFWvoM z*YAFm);ZVO^s@D_og@3AIv4eW;@DVy+4M#A)=ys)uO2jSwm78Y%1<|6q+5^jV|C|T zKV7{1i`Fk%C+}kCWdE}J$Y&q)_oMHYdrX(lIyQ>G*!tCl?Ssv4Juba|ab@Fbo6kkR zW4-tCmqPg(bcJ9dkmGV0m>*>h+d~c zTjw6trkAac?Ht(`)w!r26vxKu%ci4#Q9d@(7sZ(?`d--f#g%XA)}c80Q9W4QdDc%C zFaM(Ti`L1zXinl6^^NxJJ@a=&-tJ46$9gu3v%Yq9b*l@_6B|dzwXS|~W#ej_&qaR^ z{En4J-~30E1efQ7j?+bZhq_NY^39wlX+mzVcET3>PPis_10;Bw6FEzo5wkd zFWW!6am&^xy7%l-{$Huj(!T1fO}}Wq)}wW7?7Z3NeAp#F@~h)lI`^)pmyKipT90{s z-N$}q^K@R~>Ee4&YP+BG+P$az`a|==>dU61_3Fp!Z1pU*uNC*I%UAjyP(HT)*!AlN zSKspRFUo_~mHf!hUisk-&B=K$ibL#BXCqym&d>hnoWbvXd zW#icDWWUn9UiZ=eqB_vM7EiD3ek}IB*6#h|*B^E+)t602dDh(M;?%RK@8$oh%UAjy zP+qqF*!AlN+b3Hd_M$v!UCEF9?3JJH-mSjqW&6Ku>l;^{zZ9oG)CXG}(y@7K_g?T@ zkMd)6^RJ&SUj9Yv7p;?b(VU$3qB!&(t^LKx$42#;D_foNqde<(W#9i-+j-&I-#Lmy=UiU?#+#?Pm=`*Kb+FYdu3Ub)I-Q^W zQC;Hg&$gdth`LMcsP(NLK_pfdJqIh{1%}M-X`?qg(e@EmkyhhK9L3!Ek1G|3rqqNSs)~1){S^93PFS~u#ZXdN*zl~#;)+0X~ z>5Jmb56z9;_10niwby*)NBOY2dr&`JeD|+y{i1kz7tKlhV*9smb$>_XExl{@vA(u- z>aI;k^_bt{75C~FSK60e+k7s1=TRR0(trIOyWWLMcV6f`7n{#IIveTeT`SE4*SzUy zAN9Id7sXqzKKrA6#EWC2a}-~;UwM2d*mZ1mqIh;G|F3L5wtdyb#`d{vb@jUuU)_A2 z3mfU>^3(P8EA_1(dDLM}UupfK`uW+o_H$nsy+i0bYJ7G5FPaN~ZRd{sUn#$Hv7V0n zuasv|-sS(Q%U8BudFjrbUBCNMTG#u-UTl4A=g7XO&PDy8I5t*aHXZek^0AS=D6aGM zz0mE8E8o(sLviw>da$~4t)DJl{zdB-t&?}roWw8c8|~YB=I@BS-Ip$p^=uSpeQi3b zTU}_L*f=__b@huY8&}(WF8X_b^5~b2ufJpcUGZLgrSn4PA}^||`K_a~k&bIl=7Bwj zW%qum|DwKHr~T2sjhCP99L1OIpY8r)*Rl19;@PGAztUVv`>KnLulvc{`(QntZN6y! z=)Cx`{>F7~Uu}Ks$9nrLnqU3Z*|_5Mi{eZ7^`h_H{2EUer|(5`;;(If$dA?`9qC9% zy8N9>ZR_alMfsQit1e%u9+a2uez5CzKT7MIYi)XIo~7@$`m)<+?el2mueevgxYEA#+U9f7JCE|{mtL;lvF0v*y$hG_yinf7=Ch8@Mml=eO7p-q zZ}*DdKI%13_p7YmdiB{K?IXW9Hp(l$Y`^m82d%5EPUL5o^8d=_W7~JJeJ)#F{cglp zH(%$%MtZsYbbb9wbI^}G>TsT4Y5k)5`PsPkb6+odSA0K>ude?^bK$S;+>!q)<##UD z(~tcV9~DdVko9t&i;-*%#Hhs2>!^#_G$aqyAAoHqsZx zb-unAx_xowTe@{9PJUDmR(G!T)5XibX#Jvf@-CW__(grAeS6RR9g(;D(&e$9jpD4Y zO-FUB3(XT7N5{3UesN{vYMakRe-BU|{n8h&-?83{uXJAMT>AG)UCnPDosD!{b21O? zIV`*POZ^x1*E;Qw_HDfUbmu6(Z2xTc7rTzFPZZBC<^PrDQrcHtY;63p)!X{SS9cz0 z-l!ga5Jkz zC*Kdf`|Ar=e!4iM%a7{8>guYWE?)jc>ldw)chS3|FSH)%=16z1(7U3KQXTSG$3}6l zRG+%A`!}xhTkGU${8!t2E;<+Qzx|PZvA<*Gz4l*z=Y`H^)zLiG(b-7HUw^+k7w5@# z{^~|^DeG^Y@;Vp$na8@%;;k>+KU;pZuC_jrpIyrTE6t^}?_%SZZH}!^eD%(kjpkM^ zKfQDNYMX<8tXBvBSE_H(J?3Yl{mXt2wcT@+PaK_%&TY|r_-psP`M*+r=U_b@`Clo| zqP)xhRhO@9z4FrCGj{#G4{Yn)r`q&Ver&yLeq6e`Q9K)8^~SMFb3pyGk-jLdbMl?g zyT870<)@27y8NgftnT}+pDtehMe7%>lXuaa^o71EgP`yI=F?Z5oa3!RUB%I2|-&PF=^?eEu~kMmcz zd0f<2>y+QQxIf~1pY0>w`m+7A<;SjL>l4MZOZk7Lxs>)*7aL#qleh77_k`_y(cIDe z`LX`SInQ5d?&_4sK8xm8KixbR8^=aRwz(lcT8DI`BOU4T zo719o{Om>fm;b9SU#T9Hm+d~W>vwLYbk-&1vPK%ZKXZM}J>>Z`t0juCJ}G zSL(0v@?ztbO)uqH6whDV-&r)T(wxjq9{us7xY~5|JJvaVrE^C6^?kEn^IJz}BOQHj zrFoR*BVQ@6y2W4ApM3Ikf9Eb-3cPpO`At*_2*-miA|f3mwkON{IC67 zar~&C|N8R*{r4Z%Ub?tN`O*B{3p$&hF0cKOj`aWhQ$Nznbspxzf7v}B??dVPs9t-; znG33mZC-Swqxj3;c<19?s9n8&cImzFuKY@Kl<(rIxAloLk45(d&3D zqB%M5SGsTJiFDsT@}s%ecAsDCP`~fcy*+5A|Y zT{@3H*ZDvE9r%^bul?EAeki_nb$kB1QXHM_Ua80aC=TiNL%RI# z3(Aky>F1(%3cUyHQvTP!i~bIxb>`#y){ptuHiwIyXXE)hFMX9)96!3Bt-to_Q{1Bb zs9tlVv-#=r+8^mickhs1cFt^bd9`~!?oH_(qStoM_)Bv^b+XNij&u}%`5W(i{GF*? zy?%CSo&Fbnf7m$l;xFw(XSbf(;>=^weL?eGcJGaQ?A|U~x9VKFy7eQD&hGo`eO)?# zDZbyGzUcNxbzt+>?t9NKKdJ|-n}7Xu@$xTPzi6Gji{|9KU+KP?C(?cY$dBe;+kJkm zL;d~U*yf4urG3=PMtM;^Z0uac(b+SU2> zSIUp#(Ruqm*y`=dCd+YjmTyDum|TBo0j>%FSZmfznyePVSsibLPG ze$2nNIb7^K8_(Z)>8rfr_|g4r{k2!0;uhsc^_nA{%}En-?AFDE{&{-ud`Dv)DM_jd*_K_r2+t{mQPBcX8FDE)-`Ti|z}W z@3MPu+++84(K_{^I@w59H=Cc%?)&Ta*3$V)@%`@fMYlhy1Dm(DcOLoWNA+NJ^RJ&S zUj9Yv7p;?b(VU$3E8REqM7r-E`O(~KyU(w6sNeVJ-m%RS`R$`#Hp+|YVPoeij?P}R zf2oe@Y<{fHE}e(J6Zq@zSm)RN>}x+1U%NWL{z~~#JUVaR2V0%uoMS1D&UUZVV}BHf zbo(J)e)k3CN9*)+(L06S19mC@*S~}7@1gHoKjvTC94>aAjpy&Y^i^JQ{OEqR{@SZg zaf|Y!dd-o}=BLYRf21Sby+eB0IkU~>)$aLtcibPWUb}Jin+vLwZC-Swqxj3;c<1Bq zOzrCRvrF&J`aaEN`RU^1ySVCYed5ex(S1SlU3Tw{d+gpWTDR(4y1Mlvj?V7;>wR51 ze<{A-oxbSyM|EKH*6w@HFF&dWtDArQbn)^pTEA$Wyo=`KykF_QnJ3bH|HzN#UfX?s ztwa63L-&qtp2%+>^|Dc3R1X_FS8;UqqWw#CRA=*Jb$01Iy!++r?^y4k^J{yu@pyV`#Yc>`=dCd+YjmTyDum|TBo0j>%FSZmfyMR6RWdP z9QwZXWB#?x;bQ06c>c~yU*#3YkM3vduf6&dw^8oJEWJLGuvEV z?VgW&Q+ki8yH|8^NH-T$C)>Q}NJsIPzwyq;-t~n#o|!|b|LV2H$#-$p+xoOrSq5K``zh_ZhuqcQ&f zUq4;E{EOBvS|{(KIXUlFx^L!*bl*Snqq*01pI_@xzwgk!W1A=P+ef`@lo!>*#?Dn7 zoxN!PQXSRV{8*h`IuGyu;`KY$JLvq{pMC9z;%isu*Iy|=ibv<|`(Ue6oO3M2(b@hU zsK@>&4(awoy8P}7%8%CR=c0ECy$9@4e&?!Btj`PVjwi=AiV`8zLtl~)`; zx}UAT_UcpIqWq{{bELER>GIki=}33)kY0AqY;$?Fdp_<>={=&?cE9*bb3t{o&5Mq7 z6o2^}?|l57sa?H(c4?jd7yaGC#+es?X&*Yf_0$$;9*gb^n(wlEZ`@<|cG0?3=hD@! zA8~Yc-(T!*vCf6@9y>*QTDC+Gc2_su+! z?)yi6H22!>^J^XI@At+wPjoNsqh2=3i|S!x=PHiQUbKIyj_PcFtj;c-hrbIKuivrG zul?EAeki_nb$kB1QXHM_Ua80aC=TiNL%RI#3(Aky>F45lud1`< z_xDbpSe=dH(D$t$^RI0V7dy|!^LJkQDz7+xbU#~v?bWBaMfp*^=16Dr)8(~4(vj}o zA-(LJ+2-~uTHZMBTQT*j^yz}vQrgruE*`@Ew_g3n^dTnv? zU0n6HK5^!;=)R!&F1z=}J$7#wty^_2UETT-M`!o_^?PgS{H6GQclx5+AJu`)TiZL2 z{PLrEu)6uzPZux$qV+SU2>SIUp#(Ruqm*y|E>F0cKOj&%1E>E)WYx$s|h&!_K>dsLmxzbMXJ zP+e^Eq9Yx}U;f5BAAf&pSFfL4TG#%4Z;OpLrxkA>es=4rEzUd^-4`_9W%u5=$L{T- zb*s*$t6M+f=xpy9dXL#iUliBxPG5BUqdIWqr+fdAEq}vba^1Cl6KU$}s zi{2^p9dsMx)xKf_no=EroBR`sZZTI=L4)ynYW1A-2MRy;s%Q^6T%F?nUE$-}*8C+U9W4 z92!SAZ??J?TQ6HY8|m`@N^|ad`*@$QdTnv?F82FkW9wnt5ADY;>t~}m+IP|Y;ji63 z`O!M0BOU2Tw~zYyv3hNBrTo=@wdE_-gYvVTAG?0%Ra)ooL2Y_zo~8SC=_{`5%whGl z^!Cp#)q&n6HqsZxwSN22*{FV8`RU$Aq|1-$!RqE(KV7{1i`Fk%C-0(pEI(acxYpUX z_qKM|FPnbRIr@%Jy(q2}PuDlvpN;aOxY~602hFSWUbHUt@uRrfbaW2t`$~PI-?8dK z`!&CHbT-m)-EaFWu66R7lXJOfpVmo7dD!+dcl+`;pZL~K?>%W8+xZ+}-i`KVp>#944vUBHGH$S~p|F5)>+ky7`mO$}EnlfVl%K6`cKz;IX`Sz*HodeD+OOnCem2?{ ztFzbo>dU5M`!b)U^Oxe+cSpA`>I<8ldw)chS3~FSH)% zomamXwmwjuZ08`JEx)?dUAww<>O}jYIDM%L>EcUqzf!*!_3Qm_9sC!qM|1P;e`V`d zm*27W!REJ~&PMvh&fmFV{cQKb96CRBBR^Z6D6cu#U%dOqZ(p|crTm>ku0;);!6Hk z%8R}O>*+nO+CA6Wy@&kfT6^h@W0&h3EKYsh6da$~>>ZgmBf6@9y z>*QVZzBqquzjX8B@AtxPUFxAb2k~tA)oUKLt6Qf|6o=x}uP&skvlRC$^?Omh`feTc zi_RD2?fd^~w{CUG-`}y#YdxKf^oyOpb3=8BWA|Qke)>RuwmPx-?JwSav%mSXtuN*8 z92(a=Y;~Y@>~ihb{g-WhX`k-T_T67J$L6=Ly8LJ^*t#0WkNob>uar+e^3WGo99>_F zy?1QXU&+7d`#}4Pr?b(yEt(I1?VdM3T8DI`BOU2z&WqOZvlr!E{;#@xrFu|aw)@7e zzxRP{o%60uFXhM9%jUcQ&1 z-}>p|k-%}HNqJ<`pQ?p~pHMIWU)Qfhr#}y}@Uwpad{i>VK zMd#voQ1jCpS6dwNU;8h=^Frq!Kl&ZpJl4_KNXPYi$$6o3V|QL`aq2dYi~4Gv-H$Ea zy|y2!XR-aW<;S&NpD3PP%Kt0PrL?a)*=QaYeTR*!?i|hiV&~BK`qfRZz4XSh7n?_% zIo7|v6T1K2G1AewvyompXF9v&M}GB}=3ia^SU=l*+zVX#^7o$Ex0JW+dba)8xb*tF zf9=-6Pp?f!bs*h3tj=~0i}IH3Q$BJ0-CtaFwsmFmUVic37yo^J?EfBN=_^irzfZdT z)ydYs|Gm?q|D6MhYdv)2XRC+JPiOO&-Jief#L2s8Ue0^boY4Jro@l<_Arz0|(K_ei zy(_(c$geKzQQ!Kct3!QQ-99J|=~$gDuRa&eLq7en7aPYe)$=R$QTls__Vqi~@40`) z*6-e`OC9L^-4pBi7xlq!y>pVsxpaT=;uft(`Ac=6zSPO~{lC&a=6F$`;`E8~=~sRF z!}T3h_ubOvU%GuS>H`~3NBu5)>1!RmY#x5|t>1lEbT6=R&S&ZT==^Hakss;mFU`NY z{;__x`M4Lj^yPQY(7vU-rFEtKUUhNSb^qF}gP&fTj`l&iby%J492Vs*+oycu_`AQj z>TK)E=Dqyly*qwapxQEn6w-1U#I#y@Pt5574*z)O* zz1TQ*sh(e{k8=Ip^LzDQ|Bh|_?yb67kMnm=tZ!Ux^P^kuoaAvX-Cw-8Me9+1^Fwu@ zzSPO~{lC(F=6F$`;`E8~=~sRFL*L^?dBypD*hp93MfGFj>Bz6I+DmU78}*0wW1~Ln zckb>5(o6n}&Z*QVon7)Hzxqq_sIGsyIHa48dw|v1rFi$3ZU3_C+18cIU)}!5-~HM8 zXO~^ie&uUEbu4zj&WA0|zVx2=+J|m_rF`^d(^0(NTNl^wGwb|5LEowGl5YRfob0pM zzXRgrM|EL!_0>-oFaM(Ti`L1zXim<1(LF-<$a$jq@?-P*dyLjQ5APlN`{(_uZk_$n zJdv&*^&!3c*pH5Mtj?BSov5DDy5^;qjboSU_?6~UIv=#J^Tn^fW8I6g^}Cnqa^Kx= zwm9q6i{ja(_3G(4G_P}}m)4sX%3rEaKS+1a*Lzi+?R?DdqCUi#3#wZ`=2*M;e%a<& ze|7tst2)_8FYDKzxXXTxljov&me$do|Dt<>jdMQLo!6pu{N~55pKea2_0{!@^|N~) z+?VR?QoMV_Hs`YI`K>F(zv|+w)BmFVQQWfWXdk4@gVov2VNqVxU&)XB-QT{|+19<% zxp!ad`MtB)-(gE%apL=Z(!0O<(C@uP?>vf=A6I>J^|7s^v*lm3e$hI47rj@`d(k~X z_tW1!G-vM+ibwYVt#cmUIrI+7S6yA!qj~C|u21zKz5CdYj&!WfmS3NX<{@AEmX9CB z)uyBMzfvEidxZ9LFUqgKW4pflJAdb=4s`zF^si17&qliS&Pg6rmw5M$AH}hq8_JLR zmlyS=PPTUny$9@4e)GGiesSi4>eiq7^oQniQC@MrA2!m}cTxS=cslaytM<|x$4333 z{n)6_`klLbf%KC9qH`+sNoSY*$glp=JgV!TE)MDD;~rpjb}8QdW!t~(dbV}t@>jP% z@^^o>{@G>MvtRj|PaTWhuk&GxvoF2pz4oD-Unw7b*>n`Yes@*(?$T@feS*GI-zB~K zv&~7~#lH9AiCuBQ#v2Cuk&5(?^yRhJiBcD z?xniiclVnu&U*Etcy?*MdU_7c>zwJO_2z~0m+I3G(%tj*UR7tCxA|Svhd6UVb?e6* zYn#tSdByo&*hp93MfIb2=fb}H;?cUx7NqfA4|&Kri_(IXY-fVm)6nQ?tdxot1iwu`z+cY#VwnT_CdNlSe@-07R8m6{|tIKa@~RWnQ(7l) z$&dWC>8Orh*>hw&2ehB_z36wWdmx@&>J!z0?z{WV_WtUhztl(R+|*-#{mN@y_v5z? ztw;Gwb?OJ{?zwjgy$9^lKI*+_&f?4kou7WpqjvB8vgPGhFB|FRy5>zUyH20tF8ei3 zo>!VPTE}+&=8W``|Dy9+w4UGm*!9!RskFYje(Bfo z7xJ4!ssHNiu4lhe9<;8uIu`X)nlrs~RjwRGG@vr^FwSGFvU#e4nq`U9^ed;}^&bCkMTDCg*^^NYgICHEm z?xJ(M;@l(qsE>{GvVQ%DyX@CEc`mxwrFHb)1NVVm@?Ugri`Mg-AG?0KIhEE|*Du!3 zUiW0_?o<8l5t?&p9rDXtihI??S*QO+`=hvJ)6qUimj|n}ox`GjN^_>$ANjk#eU~oZ zEA3<7uH*MR?V{gj;%a-Z(Rb>*p{v9G>O*-Kz5k`UuHXoX^tv%f_3ZIR4t=N^#ZAAM0oP-rNgZdi~{Q6&X zFHqdF>1ZFM%Y)U~&SBBKOLL~{8~N2m$Lh7kU9|sN&+nZ@zuR!>{3s6lebUu!fA!(| z9Z2{7BV8Rx$Li{-pWb+J*!}5UCr;i)^D>8v=IFka@^((O{f_YNt6P5LM{%|3XkNe4{-yk-eo())tMjklvF3`--JHz< z)#aY)pMSA&@~hAO=E1-Alg~P|9_24vr+wsQ`~F{P{_44?PI3B1-!Z@XtV3}ZotrrI zvXL&o{n(`IC%YKcM=aud^TE}+o>wfTK;#Pin{SBk4{u2?_2_r$$e zdi~{Q6&XFHqdF>1ZFM%Y)U~&SBBKOLL~{8~N>{K3ux>ue86s zUB~bD*2VSv%sRhMu-_+L9rjlr+Glb79Z&~$eRcJS>#mT#9UOoTX=7jw2 zRk`jj-8o>-$2(S9kNoOUAGUA0c;6#dw-1U#I#y@Pt6%Kg*xr5TBR}$^xY~3yuU~2Z zQvR~Pzw}3EBfaeJSaZdmt2#Ol_e}r%i;Z)R>a)Lj@Y}!p$!8r}kMfu5lppEtd;jxQ z??H98ebjSNo#OP3?zcF7*LLm~qlJgS#9yvmtEcZa>db||Dt<> zjdMOr=Pw&?e&YCRiz~%dH-D_3?R!J_gI)5QBR?Ja7p-${?9%#IU7U6LU$j4pTQ(i- zgLHYYI@>ucnrG>|rt6!{uP%P9UR&Hn`>pl--dXg!4VTW3;;`Q*UETIqAFkhlbnidX z)q!-ZuAchojTeXApWb!i)C;?VC{b4BNF&gOvX>iyzh zG$(%f)oXwA;9vX6XB}FP`YqLg_K}zE`+sHoQP-k6`Sp#yYku`vhtBPybCVyfW261` z#m4H(Hh=5X@hevzdb##dFTcLc1L-CI#nw~X{-`c-^xE#jqWi&L+xX{L40{i_Sw_ zZTZo@`mJ8SIC&R+Cn!#SR3}!qfBkgv@-JGyXq~)^`gh)oy(ixFmAChdjrvqSo4@sx;Nmn)CF;?=7zlwZFnFF(@7^<33moh=XhE9FP$sULmUw(rIE zr%&tHXrDFT>b0E@%BPNBX%1a4e<{9HFWnrPmtHn**{!Fx{ZU=w=(XLGMfZolcJC*D z`>;RqU$!_DuMe!wR?kfd=U_MZ4oHLv=`%SWR)_0#;;A-jER zw_ZQ$VlUsVU9`jH3OUuEy}t#+@lb*jhwoiDC_PyVhm2ft&@ z7wzAA)VnCYes+D!uMV`ocx*oL?nx*xLdn?rTL?LbkC)J{V!V2PoF3r#nV`O>*g=57hg8Me%d_5 z)4p@)dr2evqB!>$?L*ri^@l4zyZ^l$vdfR^!P?!g`q{zzj_o?X2C)LlFOvfE!j$nJaKJB7XnbZI~N zp3?rMe$1nq{Y2|les+1xw{#xv2|vxh?CP*j^|Fgw?0)KDFYDL;iOyA>xYoHBi|xn! z#Pe6L{n%SqwLXyF{-ycVPS>wqux7Om%}aoPFN z{xsHJEv~e$__Fc!)8-+b_MJuFR~p$D#kt35AKLz?KV13Q`@LtEAJv1kyI=LQiR1drR zC(G7(Y+T}?H!il*;l>n@-Cj< zrFD%{2Ya=3>?po;Znew9E>6DET-cpAT953=Zhq*ViNoGU_rZIkt=A8#Q-8JF-@8J6 z`W}l@*P{H`ec9E8;;Py8ccS^eEA5N=_1%9;=Y`HyesiJC8=Z4??flNMx^qzn`W@@s zQUBIECsaSyu5We8gZj2!J>6HlIQyXdr9KwjclmE=zEAIicylTBqpoW6Inn(3Zme5& z>%?L0)t#Gl&i9tiS9~?Q{AK&1)k~W%x*v4u98i4OIjFaGx_%nf?cOZ9NBq_O?qbIyz@tMEcub&J-cQ5qi<;r z;%V>ZMDJZ(b>Af|9?e%BG|Ic^{gvWtr|nNSp2o&ohwI#|ZydY4C4ckEgZ%CjyZNDe z>b_y^&cS^^@$#cOtJ$qrhy3mXiqp^Hnp5r7>S3qR98NT6_30bg%dX?Mj=$gkTix&D zIxqVjRbFX$j*72kI z(S4tCy}{pQUxZ zhx9F*m+pG@zRQzdeJy+Yr%Uezy(3(9eiVn^UHz+GcJFJk?^K+9=(7H$+4a#pbnBq4 zZydY4C4cXM?+o(0SM278?zMQdj&>isH`+bb2dY!QwReAYxDO~!9gFf~_hl~|N29t< zG~aineNn%@`%me-(7DQQ9<+14rSmzlb5RF6Z}dA>|JJ)#sD7+n-_B1S)VKBO>Aox8 ze)4w?`YF9@c6Hp+K2Pt1cylSuL4EpLTyeEGFD-BTHg6Qa;`q@z`EP0ciq}``sFsJF z#&y5gOZ~hnn^%3!%P!viUi5c>zuJ98{xxTI^D=j2NAb%q4#l@m{bIidCwu++EITKD z`(f?PFAjT-)z-7qi@r~2z9m1_zqEK%SLvLPpKkr^;>_{Ho=dekqkZke?w+Zmly}iP zDO-Qn^Rugm=4Yq7Pqp=n;^kd5FZaQ@SF(6znvWgrV}0p6 zgs=LolSg0dC#qB3jb}em9u)6>JZ0~jyo=7~mhE$Cb6tM*^ggWLvEoW~vwNTRM|C)7 z`+F}mitqZ9T^;fuyYuZoFF5DgY5VqkmzIy;Tu?ponq#&26ZK_3tX&@GP&yBD=cnsu z7iZqD(p=OnkA2Xb78_^2@*_L4H!pjs@6w!Vr|VZQyS^9I!C!5D$lpFs_T{%Pc3!mo zkX_zId6)mGUA|IXC@-yEx_;+YTIam0*~`{PJ4YI8uNGIfkLuNLk-PjgaF_mQu5e^-2`&^@M`m&U%ArPZ$< z8s)!S0Px7zR6 zwO$|mw0!F0$DX@*b4U5am-6eImhVLC(0aP$f0e5qb6>pOPySVR?e;U*>XY4hqV?}e z`|9VEPh8nN>m9PAd*>b^`-#od`1)xyPyhV0==;cD?Vr1lzxSzg<7Z!*9mV&&YNyq+ z*uGZWtzEvd??OIy{n7R7XVE<5p%bq>6 zotM0e=7!?zi|wa&^@-zWr{!O?e$hI47kyV&9(MUqot-Of9a=}b7v39fzRs^yFY?>B z=c-Qqi&Fv!nc{zQvdJMSaty{hryNHqKa1ue553qt)|D%+_U^;azWQP>+yBz;F?x?Q zvM-A3JMuo+TW9CB^0S)^-S@h* zIjR%I)7buK`_mj+v~}Ja?f$l2b(#-9yK}T2#iO{^qb_#w<(l`c-To)qSD($#f1>p$kMI6dIxjRg z=YaMTcT4x+MEjz1Lw%O@s~1ndWA)+u>~9}hzP=Z6C(7SBoYwK{6V*kR_IY|A#G6ZL zUv*V`rze`DxuA74+NW$C%P$Vab#70!Jo-7U7gsjV?RP+&_lf!}`Jd80-huV(?%Ses zPPu#_0sjXU)s9fLweCVv~S6e{4}Zu zYqzd_EqmkWvi;eA+4)QH-kta1-O$}%{n)(K=7Rk8MfG6q>ZzYyy!?ySFIp$>qVI}$ zY+iQ#^LMVaK2SU@&U>TfcYf-tuH89WkM=`xt%sj|%}d;?)aQxz*C+jy)}y-n{lC@D z3(d>0P6Iv_9$j^|PpNdFWf7 z^2#HQM(+S?UwyHc?SEWVcT#&#P3|iRyW}&swJs ze&7A4bY6IxmvgP?JOa5ER zkKU8@?5FR6a}-zId&b}UKx5ZDI`=7RR4&5Iq`Q9QC^?X>e*b=JQ6Wk=r= zl!smvw`LVlK$nUc;WAC;Bn}YU@vwZ^iN3-<(YKiz!n=sR}e`tDih`;XqKcff96R4>|Ras9ne2X=k!>Ji7!POtsg ztzQ%;@1lA2ysFI!`Q5j2-CuU+fbz=I?-Ff2viBXSuYI$lby&N7P#m&j?XwM%#eiT>Dj^-uqDXmBK$Y1K`mcCPq@|Y{SKjy3s{n6&Ge*VSA*{{E2&7u2uKl$WA z>rwyurPZk)WOv{FUPa#nx|HAkC+b6-K2YEMji=4)MCU6%T1TV(&5y=qufK8Xc$Mi>8v3;$$Tf2Ov z_lNS)`lRdEAGS|g9{QH@mR-+ZyL*NF^y-VfZ2wEMpXfZqRd;`L)NlQ?IFxtMJ3(>s zp*pd4`_<1bUj9Yv7p;?b(VUz&`VJwx`@ybWw2!#1V^_C2=~a()_HEwouMYN79)Guy z9qmJx{3uT~ySh%)uX$LH?61;2eoE(sty4YvbFR1SIW27tXx?c5*2AyvQv9;(TYkS| z)!}}uJn}YPT_}J1mzN*e#rIs*T|2$j@t-JvsULk;tN+CIr(Wx5w2$+kaoPFT`KaSn znnTyKmyM&km7Z+LvZWdDM@!)9P7lUn}m` zE??ubMVe_Gc(OE-_Y#Pe78{#3g^?A8AF49IW4MejvDD4w6LpT-rZ zo~QIq%Tpe6DpwwM@1C{~Z62uK^5hrqexNy){K!wE_k^{#Kl+yDAikP?vGLXBhUSFq zXdfEoU39MHDc*kKmS)H9bJFrRPCk7uI)C){71`Yn_U1Epbnj?Xm%33r+ON8H>(wJa zjpFo!^3jX>#>Ul7qyBHH?(W0>Dt-5#vfoF2tU8>-Emxk}&C^_*AF9i}p~chnvt#}G z=a&!tj#Y2#>ALnIk2ra-^@w8^r_Nhi_w+u9?>>EJ=Ab{i->aqN>ptRZHy3@;$X@a< z-rnOCXFvUjW2aASym`6@$d2sJ0ohCI7Ols|Eql+ay60N$-XVW!&d7h#;!u3grFPo7 zMg1=SQ@ebndQe_kKXm=>S!tbjQq5l42klq#BR`Gy#oFn$zV@ZrvF}8m%g$elKfN3K zqQ0%0Gyty^94m(6cIJB{oocK*%{)g_L0 zFKF{?-E{rx#OAlZ-?8qS{mh@XzLdXnXk7Eq>Okx0a_!gsm$tsNPxq(I1I?Sp+V#m! zi(mWJt`4kUUVi#1)ny%x?29Xo-92-^k-g-9S2nNtH7~n(-}}YJRrj9pyAO-z<$WVN zc0K!|c<0ytZ?!p<`deIc?7r5s+lQ{-yMETR^3$!KU0mWLsgViPy1`Ub5=LMI_&q9?w$O|j_l3PUh2DakJ(RJf8zALs1E*W z^F#ipl)v*~NB&!$>X7#ock7q0Y`yZay9adr&aJf0`B$@-&Z~9a+E-ljwU6r6Z{z5) z`>ndz7v*Vw^Q@g#Kd$`jzKh5%AF2myw_pA2;^kkoe$hI47f*9iPxq0pc7IoVr+Rldx(uQs0(&B^@aF$eY&8!xZ>@svGp``f2%e(Tw3WJm9< zG!N{&owvNz=3srP4|ZBU_3@)R8gDoSqDpx(`{;o8i?jvvQ z=DY4^?bXf;^+f3{!)CYUUqrR1KE+?I%F^Hv)FpgOFVyd?_afh$X?y=C%-;X ze^~p{>}b7yuy$HKPuY7+pYoVfX+6qAJ8!yv{iFWMlV7~JQXRD~E#IQ^kyn58V&j%( z$Ih*G{!(1)(O2zs=hD2@y%+rQp?a`(`_|7cUj9Yv7p;?b(VU$3qB!(DQZK4UoP0Ej zQzy+YKgzSdkJjr`J*7B#b`@uzjJWDxcWZ% zyUrZYe6i%?DF~@Yn^=JI|p@O^U?14E$#F4 zK8QCLw6DJOt3TA|iTY9}T1O+hbu`vqZN6xoK5nT`@zw0|m-;Ey%P!BNJ{DJ9E1q2) zXn%fW7w6tAx<~xg{&U;NZyt-jLu=0b){CoNc5&)?O1~e=Qyz0FR~~lnl(r9T9;mc77Cx?Tg)f)j^~AE{ekySGzjdmu5%trTt6m zmfai{J2&4O9twSi5s@A5gsfi~ZhOuO9i;tA26xqQARnKYFooG@8SS>R$6; zKT#eO@4Nq${XXhL{?a+z()pd(d(nOC?>URFy;>f2{p*__owNLC?zDXFQGdtE-#OR^ z<*_f_b+`KIeGqRh=$+|LeKeZOiRNZ5XdR90*3npdwK<}7=6*|ki?3#vztm5uUUqrp zNA^-*D8A%hR0rCp@$BN$O|QcMk3YikE-Ucb4CJb;z%76h|-m`-=9X7aK>T zIh^QR);!oxln2H8?muO}kNS|mbPl(4ekVG2?7r;gPCIY$^kUcPU*G&_Kl#zzY5R$H zpZt!MzjF|8KV+A;>u&Ya`yk$2&^y!DqPkCPpN;3YkNKH9jqGLp&P&`$-^R&vOXu2n zcJuB#V=o)G^y%Joo%vbEzcf2Go?TvKFWm!lPw0~09FZN_7v*(6D{lFhT^@e@FS>sy zZfSP353(<={OqN9mfkJ9c_P2M%mtU7|CaV!`S^Y3PVD!w+CML#cj_ImtKa_WLwOhd z{|_%)f7kP~tB2-ir@K$J^^4-=T{JIqFsEvEwgW`}KYp3PaA9ilE`|EtG*@(g%vnA9q|IIZ{ELmV zpZaL$(*3)ieAc1$=y$ArP@VcgcK7|4zf6F>2XyIt?0=#@#OVX|&EI(1z9-5j&ODY~ z{_bBpKi01we&@?xz3h#n(YnPI$L=01x)<0u_h;Gp%f_3ZIR0vJrMTM7AM2-m2hjbX zOMY|YXGi`;>zp55T7PR7XPy2R?T_M?W=H!VyF6Gs?Hm@(v-Gan^-c4uiyv#R7I&ik z)_Q)w)3Cq8mVL#E@B3s|xBbWXIZRdCdvUtMuJhxBSSD;;Pxv zyk4dKOZm(8NiUiMKV6zDI(Kt62UJ(@7yqI;@yoA1`dWH$S?5cJnN)uU)_F;*j0@b{}xrm%sO{^QYy-Wk30E*O^19|JrH!tJzPq z4#iiiW6``z^JX_Mn%}n6*!Rh9|5aaq4_5c@0KYnr z9cx!l{p^hwhw7lQ>%_^sXkN~H(K(`f>+c1cukQ->K8tfM^3uL@rF_V*Zh27O`eavs z>#E&8*m!p2ua;MTXbz>meE;P`eiT>Dj@G|Q{gvj6_S0AIM|I=)7rVdnH)s2!x}3lM z`4>A+`PHQk`S{m<;;ch;(mv`rQJv!SjqW$U`mB4(=4n0}<+qN; z+Nh^bQv4791d$0M;W6}Fr zbLO{RT=lYxQ_oZSdsUwDtU0bc?7kPYeQ5JQ{go%bc;}DiSn?x(@7qbYKl+yDAfEPa zPW0ZzReO(RTvJ{z~z+)9R)hPh;b)!*yQPH;!H2lD~Q7L4Nm%-Tcr!b>Fad z=iok|c==JC)$G=*L;l_iar#+YbE>^sJ?u1^!-?jsK7AwqiPod}_WOUU`<>K>y3l^k z;g-(tMEhXp(7(g9y2Lw|QeJji|EurD+sE8#b%=Lg+*kQK2ft(WQ(7SQ<1Md!dzqw^_yf7Dz1 zNn2;0?#ZIxQT)}tfBfdL=sUFL%x}H8>SY(Fp2haH;%@EomEIxBN9&iaUq6fHAP>Fh z-r_0FeGzxMXZ+T!zSvLfJZSF-y(3(9eiT>wPKYyK_0lMQ(fh*{SG)bYuG%{BbZK2_ zUG3KMm;BBf-2-GtcJo8`%>BgLorC*;;#ME_cHL$ar#)47rP&O=P8bzM*W?r zF8!ALCw86l==c9t_xq?1{h@uG1Dbzz?fkx*)$*ghuz!c?&dWKL`e3K^uP*zeI8+a= z{oNOoztqR+cdY$yX}-7jpmFAc-m`wZ!)o{EMDy!BtXp>L#L>9yJtuL_`IhD{zM5VB zvVGC&r#n~k#`%~j-d>ZHe+;a8Beq!fA_nj>5o$*)uo};+FOY>DPjq)za zSBkHlZas~svGMYt??7pN&W6`@> zbLO{RT=lYxQ_oYb@8l_uIhD;P54(3y*KZ!Ezw+c4@BGmmOMdr@-#t3fy7ouk(p<#T z-VJ)kxa|BW?nLud2aUJ$Fqcz2zdEg}ZX7>*DNkwLvh$;P(Oq}im*4!*JuBT`cIV(e zAiw;m&T4k+)giw=eZR%E-)ig7J~YaU;;Py8cVgeKzOnVvUB_-6f4~2?y5C2A$X_}S zbS~Ak^Sd|I-A7;8zr)VMJj6MdTk2n3twSDm`uo<{akp4w^tlK*y|;@pF!+0|Wo7o~dG-J{aH*{j`G_ELS=e%Keq zqx?^MEeV40VcJF&pK3s9NtGDZ_trJg|)|J-PZasg= z@4T`1gB|(J58X3y*!$={cyF}zzIUiE-^tqT?_HsLf#TG)C_i>zc6Fh+YIgmdXuj6b z$bOZ?*G6^OIlSw0Wa*uCATmIaccf-7oAtcmC?NPQ9i2*=c>NOCEF{)~Toa zikDYB%3tbBKgjNPtUR|g-`ji8ICClO?>kfN+)q?r^IKPY=SF*Xbje?etNo;%zqo34 z{gwJF)x+*SxqsE{$dBx$d;667l;3*xU-|N)F1la*)&KI#^~m3QU_Quy(&AA3Iv;j% z>RD`GEAG}VU#WhSk5(^TzrITAdJpME>(IU>zkT^v-rBpaeaTn*((LFvhQ4o4X+8RT z;2qNTMfIca|DwMWC~nouudb!77l-yi>#NyIaq^*gl+GL1eW>3&(L8;ZQ9fF{_eQ(7 z{q9x!9Tp2g*~;ZvPYQD-VtACz=mh@4Nq$&I|P=zjsHAdzJ2m z^|XD_xuH7HJ#c@`L7n1|yeD^YlK5H~ZhuwVLgKBoxieM^4( z@~^zLcU}9EulA+c(RUX8j(bY$vEOsrzNmie@4xDP|M^$F{OVfTdU3e+W0!AHoP1~= zrSryhAL=(xG*91Ul#dqgz0vM%zkAibhwS#X9^0q5)}te@9!Ti?p3-M*3svX_IIp(p56!X=2Gg*9IMUeMCaCbV_og$U7CCCG=Itel=4_t&Fo!e77XZfsWci$GBCx3PC8Nd60=7Y5_&5qWu^I;dKp2haH;%@EomFh?N zX!X+d>#MY`_mEz+4((g=BY*2SY3th8vNw({eP_|{xTmxpy-#HC{_4l|JCNPq17tr@ zKjM~VNBi7T{>JP7#NLD6KYsU#zjLL1r%^mD&U>Tf?{}}-_mJH=T95j6AH=mDb+P-- zmf~>v*~OtgPqbd2G_s#)K4^Wv|F_zCp}yodXS#WU4NziO7*atuX|9*Lz4WT8H*6`R&WU^48vU?MuGemu5%bS@gT@ zDXqtT&uROj`qA&d#r{1Ix9a6r*V5LDL;G2W%g&GD)FY33$S$tzIqI{NS3k(^oqPY~`W`c$7=IA(Yf{ASXaAyP&$X&Y5wx$Z#?asmuA;rslQS^?B?qpRI?*LvX{>5DfKPC z_3XVT)$SvEwR?yBy$9xl{3k6A#jo>W7pI=Z_O;?}?edlCNBL;=()H`Bw66D%UbGJF zTk_kNf90*c>)MxmwJ*(%zO(3e+*4YQ{hrhIMfIcKeT)4&Aa2#mudb!77l-z<4ws!D z#mR@}Q95s2_o06CMDz4rM)_#*-W%=S_PbZ@d&q8I>#==`Ydz{>_nkda9VkyVyZukJ zuRJuepJ+a4z3={0Ixp0h{QmyY;$EeDVLfeMbZ)2)bPwELb5N%^WG`18@~KB2_mEv& z*>luqDX)Hz-8=XG7yCQbK2Pt1cylTBWscS6bE0$WyRoiz_n>qRwbT6N$=`U|IWNtw zzfymtdf3g^J*Z|!eq=A5*Hh|Se(TwLPpaKV_GZR+~S7}}EA-!lF+PCDlFaOG0d)Kuu`D$O99erof@3^P59{W9~?ThM1 zzxx*ZcR<{#mtS2=TQ3gnXB{p(KZ=tN&7*YQxb8#!=85L%yNvSD;=MQ8z3q3e+V_y% zzSd*=6xVvx#qK+MqB>BXYIggdXkU3~WIxe-(0bqfr*vMZFZuobqs6^S_riMGzUbUg z9q1mozviG$amZe-I^dPFfeg96> z*NWr!{j;CB(8yla-}#Df9F3phP#4O(==+Z1eGCEcdM7ZaWu-0 z>M>7tnxEai&KcQ}-M+|PntN$3>?iG>igT|Q-5*?W-i0`GLHp6>#g6PK9@(*W+WDx* z_l8DxTz3A|Kf5_p%eyE)zxA~5slFD)na`qmm*&UrULn8x>OSGJtE+4u{OsmZJKgW6 ze{NrP{!)D3oqpKOQ+;LgR{NhbAiw;m9<1GWq<(hs@-JGyXq~)^=HwpS(!DZQbk9nD zG=KL2-QU|f)Zh2!{+TDT+eaTX%8Tluv2zv2PNV$TInn&;EBTSXnjPgguhRDm`K#HT zr~c4>&iO>=FQ0h2lppnl&fEK-8}A%TeX!H!pdR^99I~qm+2wa%Q2tUsC;CpI?*U!f z$M0Bw4_~EsrXTaKc3)34FLCBbBm0S`cz!gWo>#U0YG0aNepG+o0nIPozUGYV$ZlU` zFU`F)7xt6(9>jS^z9ZG_$iFDgT+n{Bd9foqibr;=opwIz@eXKY$7SbVeY1O~)$%UN z&u=~LyR5H8aptpV-lh4myI08XzPeAi?CL7p2S2;H)K2&N>-)Ow{H6H5JN>Ymr~1m~ zt?u`pUw%{%*6uq}Kf8GO7p-5kPToaxau06lUYRSpXC*(Hzx#mh?`<9G?|XCq%oEw| zqYoP8MfK3wxr$?_QGV>4XnysT{K#L;j`EwAdC=v_?>zN~_H)iBI)C}Z)1~~VFLd7C z2i5A>+e|iMnC3X?Y^F9UgFG= zM)ng=@%(5$J+Es0)xI>l{HXrE1DapFea#u!k=?$?UYdJpF6<}mJ&5y;d~d4Rk$+K~ zxuE@M^I}JK6p!p!JMDbb;~mh*j?2!!`eyfTtL0skpWk}gcUfPH;>>5!yi4Sq@(|DyGa z*2%kQPVT`i-79lN_pIbc^LHQ6{k^S2{e5rlpLrs?ee^-2yr>=;J6Cb+G|G>i6V0!_ zk{|i2*-?J;G7q{u`JJc!(0AHMoK*1ge>`B%HICz_WybEJ{|#8W&!norNGT7R`K%`QKx zzwdzN7jIv4Ms{SkFS3{BUYZO0NqZ0C{7zczcZhDhd&+MvXg}J#*pVH@BRkekJ0JCU z2Q;$dvh%m^lkQyPFXd;Km-b!O*P=M{Sv2p`{Mg+qqB*$-w{)+}72UIvAI;x= zK==2y4)yoFxqs$~?Do+Ijq;*;XzX0YvC}9&c1|?E`bvJ}uVzR2&8zhNLjG!Y=czxm zpL0IZ`O7DsF6Bpkq4V}W=*Bz8QXlNJIjBcI6o>5ULU#Gx7nHx$&xz}MRXc4TfA`9# zzhk{K{g{8X`+A~zi8DtU*-t#h^P~CnysGtA`_k<4qx$;}XnyhbHD_c;cKafGY3`-D zu%EQ|AkI7T9jRtV{zY--g7%}$iyhffJhEf$wDVDqcR(XME<69~o4xOqmUmHpe)(wM zWqmD*GoMBCF3peKy+VHX)qTQcS6A6S_}R^+cDmnRzqgj1zZBnhryq9nRA1S=)xPt{ zFF&dWYxfiG{5>ve&nxaNBPamJm~V|cb@t~`#I+moxgnI=~8~w7dmh6gKoTY zEcL-o`@5hX`A{6Ps|(rXcVAHcQa>m9PNDAsUD`)~i?_dH-5dRwf3-QB*m{q=oacK%X)-<`hL?T_lf=B@7co?m`c57ut}^|OnYf6@9y>*QTD zC+B@j_scxdJ@o#OAI%-z6YFp5P=DVWZJx+(AAQg$FRF*e&Q%;cjrK$PEs84}&(BVm z&ZFNky7WC`r_p)p5ADl;qI;k|@pLIa*6+N%4_bZVontAEoi+#S?T_M+T|BbO@4le8 zQa>lI?^W%zef%9fvA<)@1-)nen18kVdZKxWQ!kC|C!XT@(S7WBRqL?iF#h;zT(m#4HItv466A8lUj$d2NX9c!nZkMErj8+oBL;;$Zj8f&?qmehsMrT96OEjW9LNktFPoo{%UrVzuzgkG}qc`=czxmpL0IZ z`O7DsF6Bpkq4V}W=*Bz8QXlMe`<4&IA-lSeU4Hik+b^x36Md)9_kb?t_j}-$>vybo zrXTaKHjfkCYjNtOk^RI|JU_aR&ege9FMH!?lpodaU9i*q?DjQZWJh-UB714>rMa-5 zwEHg3JL>m^F2yg3Hy5-YZC>oij^dFWYp0!$I@LuZJ1#rF_b3k87p*U?XQ$<%y<>eX ziZh=@^DfPg-MvD7_tkyEWmi|(KKR+qrFPo)41JGjWM355cc&kA^Hg8C^0WK?BfI>l z9<1GWq<(hs@-JGyXq~)^=HwpS(!DZQbk9nDG=KL2-QU|f)Zh2!{+TDT+eaTX%8Tlu zv2zv2PNV$TInn&;EBTSXnjPgg2Xm*(lizvj5AEliPqdGG;^|U;)E7E$?}Ki--+Ivxc-iHkMv{y)$Z$w<|WP?X=Fd~ z6wi<5)AOpr3m|X?bYhWqmD*GoMBCF3peKy+VHX)qTQc zS6A6S_}R^+cDmnR-`8d5FU9xW>4)7s)mJueb-(xg@}qjNcHfcu*~QDhX#Jvf@-CW_ zdvHtl%3RSsEBVp<-3N4kZ|hKh-<$hqp2%(=eb6W`s)xqTRUA8w@?+;j^Q*7qNB(Md zl;6Bd^FjWlou~fLe$M$s`^YDrF6Bpkq4V}W=*Bz8QXlNJ@4I^BLvhHiE@YSAeL?w4 z{hYYISGCjj(cg*Q8;bWk)_u{B`ByvV6V)fqeW8)voKKV&#d}9IKfCj*UiQY(sDArk z?Ms{MiO#{j!?jL7rMlTq+ItYEkFtI9J6B}CrM##0{H67!{n=^nq3nC%U-tT2SG9MC z{7bW=I?#H3VePazELvY$&))U??vHng%f9+zNBvf_+n?_D*YB-m=P$*3m*`#6_Frt? zYTtR}mmk%Gwfj!g&n{m6Me7%>lXuaadj8e!3GzE{G+%Mp`{*8cpXmPj-j%+W?CQ6F zX`X0bb?F1yyN~_YksWKN<<}pYS1I25k{|i2*-?Kzk6W#7v`+tMKl47(`8z-Hbg56& z7dmhEo>s3o_ovhcJKZ|;iRuyG{p>5wI`P+Q3sPL#j#?Bd+l zMfZch+Pz19^H}um)Pdsp>H2A0aq4+W@1;EDF{g6nVej|HKK$l^`Ylg>@x2%3Sc>bt z>HS>Vy7ouk(j3Io>cdrk?bX(!x@k0Db9)g8l8jtfa2}5Xn%g|)g!-t8%Hnto}>L}e}|A=96QagzgwDD>tsK%`-u15 zf69Iz5Y&en;yue|O_DZaEH%41*J zyq@yueGqRh=so%!tG?Sl+qb#3A92fWoj9z$+PR^1=6lPYQ?)$um+gyIFTLi=UaA9o zF0}p^)q(nHJi9peX3;(3ulCPr$ZsBtzB6mi{ML)BUUqTnc}l+@%Tpe6DpwwM@07L= zZ62t<^5hrq{Lvgse&p}HI_dUD-_jh!)9&+$>JwM({h@bCqxq_XMtK*_u{_1Ex~=0c zyU$6>*Eo62q2#w7*^ym-cJo8`)P2L+orC*;;^kk|2fy{|kY5~%)6b%Ngx1rGjib>V zPBdrKC$g93SGzv>efOWT-$#AOUpj|dI=>Um0o$+h@pqV3pLpwUsegTUopbI!=5Akp z@#<5r{G9{8cx0E)ez&yG)B7OaT+lmHpWm_SMDsY&yxNa-%Wj=GtbJ+cYn}PtQvc$s z+2t?W7p-1e9{G{I)E9~`?X%c=sH_YC>XW6`@(2a4yX>!)$Wspl!z zck-0SoXX~thuyoU>o*V7Z+Y^I@4Ya`a>cnvw`_m(EzLzd-TIc+7k_p8@J?$tUv?g{D;(hm@vfoGf&p7| zukYrypZxNPXO~yJ`-0*-2l4hrc6qz*RzJND;>`uUGyPzH$C}59oum0!N0-ZAd$s*g zUh}@Ca}i(7E`PbsO}+f`F4||&`Ez3UdO6$r~9KZQt_ov-o+Bvun$S?n5 z&)Ir)$Zx%PaqYLd-*M|}_uWTv)$IB^vFEN2Jgw(%9(LdTr|kF9xpsc`bsp$ks%z($ zzuFwoe%Ny|e|3tpz7)q!>t9{`xb`u3`^u+I_XXuI#iRYKqwRM~^L=_B#G4CxXX+~b zj{W$ffAI_7`ad82<9naI$@J#mAAj`oZ$|kG+7|YPzxv)Ey!XkcKl%K3{^O5-`spX% z{^^fC`^l$2{;MDS;lF+F&)?x+-}}Y475x3*{HK5acA{_py({?j_rCqr3cmHp4?q9m zw}0}Z&p!R(Uw-;ufB5;EYX0KWPk!)+zxAY`_uv2E`yYJp-QWH0_kQn#-~Zl+AAaY9 z_dop3n}5Ck!3Q6H_nn46TlBa7%YXjMCirjv?jQd3+spa(-~Z}stN-sC^_2ww z;jcgc(T_j-l!@V|cYOJ807<;i#b zx4tr+U;o;M1f_VGvG{N`_ex$NKk{O|m80uL4c!|azIbDQ b^wXc|x9*QW`oKa2e*NdmcR literal 0 HcmV?d00001 diff --git a/mkdir b/mkdir new file mode 100644 index 0000000..e69de29 diff --git a/scripts/exp_sir_profile.py b/scripts/exp_sir_profile.py new file mode 100644 index 0000000..9bdea03 --- /dev/null +++ b/scripts/exp_sir_profile.py @@ -0,0 +1,112 @@ + +import os +import sys +from typing import List, Optional + +ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +if ROOT not in sys.path: + sys.path.insert(0, ROOT) + +import argparse +import torch + +from inc.test_ms import Tester +try: + from inc.header import SIR_STATES as _SIR_STATES + SIR_STATES = _SIR_STATES +except Exception: + class _SIR: + S, I, R = 0, 1, 2 + SIR_STATES = _SIR() + +def parse_obs_ts(s: Optional[str]) -> Optional[List[int]]: + if s is None: + return None + s = s.strip() + if not s: + return None + ts = [int(x) for x in s.split(",") if x.strip() != ""] + ts = sorted(set(ts)) + return ts if len(ts) > 0 else None + + +def main_factory(obs_ts: Optional[List[int]]): + def main(data): + # data.y: (nodes, T+1) + y0 = data.y + y = y0.detach().to("cpu", dtype=torch.long) + + T = int(data.T.item()) if torch.is_tensor(data.T) else int(data.T) + n = int(data.num_nodes) + + assert y.dim() == 2 and y.size(0) == n and y.size(1) == T + 1, \ + f"Expect y shape (nodes, T+1)=({n},{T+1}), got {tuple(y.shape)}" + + obs_set = set(obs_ts or []) + obs_set.add(T) + + def ratio_at(t: int): + yt = y[:, t] + cS = int((yt == SIR_STATES.S).sum().item()) + cI = int((yt == SIR_STATES.I).sum().item()) + cR = int((yt == SIR_STATES.R).sum().item()) + tot = cS + cI + cR + if tot == 0: + return (0, 0, 0, 0.0, 0.0, 0.0) + return (cS, cI, cR, cS / tot, cI / tot, cR / tot) + + print("=" * 80, flush=True) + print(f"[SIR PROFILE] nodes={n}, T={T}", flush=True) + print(f"[SIR PROFILE] observed ts = {sorted(obs_set)}", flush=True) + print("-" * 80, flush=True) + print(f"{'t':>3} {'tag':>6} {'S%':>8} {'I%':>8} {'R%':>8} {'(S,I,R counts)':>20}", flush=True) + + # per-time + for t in range(T + 1): + cS, cI, cR, rS, rI, rR = ratio_at(t) + tag = "OBS" if t in obs_set else "UNOBS" + print(f"{t:>3} {tag:>6} {rS:>8.4f} {rI:>8.4f} {rR:>8.4f} ({cS},{cI},{cR})", flush=True) + + # aggregate on unobserved + unobs_ts = [t for t in range(T + 1) if t not in obs_set] + if len(unobs_ts) == 0: + print("-" * 80, flush=True) + print("[SIR PROFILE] No unobserved time steps under current obs_ts.", flush=True) + print("=" * 80, flush=True) + return y0 + + cS_all = cI_all = cR_all = 0 + for t in unobs_ts: + cS, cI, cR, *_ = ratio_at(t) + cS_all += cS + cI_all += cI + cR_all += cR + tot_all = cS_all + cI_all + cR_all + rS_all = cS_all / tot_all + rI_all = cI_all / tot_all + rR_all = cR_all / tot_all + + print("-" * 80, flush=True) + print(f"[SIR PROFILE] UNOBS ts = {unobs_ts}", flush=True) + print(f"[SIR PROFILE] UNOBS weighted ratio: S={rS_all:.4f}, I={rI_all:.4f}, R={rR_all:.4f}", flush=True) + print("=" * 80, flush=True) + + return y0 + + return main + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--dataset", required=True) + ap.add_argument("--seed", type=int, default=0) + ap.add_argument("--data_dir", type=str, default="input") + ap.add_argument("--device", type=str, default="cpu") + ap.add_argument("--obs_ts", type=str, default=None, help='e.g. "0,3,5,9"') + args = ap.parse_args() + + obs_ts = parse_obs_ts(args.obs_ts) + device = torch.device(args.device) + + tester = Tester(args.data_dir, device, main_factory(obs_ts)) + tester.test(datasets=[args.dataset], seed=args.seed, rep=1) diff --git a/tests/test_ms_selfcheck.py b/tests/test_ms_selfcheck.py deleted file mode 100644 index 29553a1..0000000 --- a/tests/test_ms_selfcheck.py +++ /dev/null @@ -1,49 +0,0 @@ -# tests/test_ms_selfcheck.py -import torch -from types import SimpleNamespace -from inc.samp_ms import samp_ms -import collections - -import numpy as np - -class DummyQNet: - def __init__(self, edge_index, T, device="cpu"): - self.T = int(T) - self.n_nodes = int(edge_index.max().item()) + 1 - self.device = torch.device(device) - neigh = [[] for _ in range(self.n_nodes)] - for u, v in edge_index.t().tolist(): - neigh[u].append(v); neigh[v].append(u) - self.neighbs = [torch.tensor(n, device=self.device, dtype=torch.long) for n in neigh] - deg = torch.tensor([len(n) for n in neigh], device=self.device).long() - self.rem = (deg + 1).view(-1, 1) # (nodes,1) - self.n_inf = torch.full((self.n_nodes, 1), self.n_nodes + 2, device=self.device) - self.zero = torch.tensor(0.0, device=self.device) - - def forward(self, yT, orig=True): - T, n, dev = self.T, self.n_nodes, self.device - zI = torch.zeros(T, n, 1, device=dev) - zR = torch.zeros(T, n, 1, device=dev) - return zI, zR, zI, zR - -def test_selfcheck(): - edge_index = torch.tensor([[0,1,1,2,2,3,3,4], - [1,0,2,1,3,2,4,3]], dtype=torch.long) - T, n = 6, 5 - q_net = DummyQNet(edge_index, T) - y3 = torch.tensor([0,1,1,0,0], dtype=torch.long) - y5 = torch.tensor([0,2,1,1,0], dtype=torch.long) - y_T = y5.clone() - obs_times = [3, 5, T] - obs_states = torch.stack([y3, y5, y_T]) - - zI0, zR0, zI, zR = q_net.forward(y_T, orig=True) - - Y, lq = samp_ms(q_net, y_T, zI, zR, n_samples=8, compute_lik=True, - obs_times=obs_times, obs_states=obs_states) - - assert Y.shape == (T, n, 8) - assert lq.shape == (8,) - assert (Y[1:] >= Y[:-1]).all().item() - assert torch.equal(Y[3, :, 0], y3) - assert torch.equal(Y[5, :, 0], y5) From 74bcf02891dbe9c467a79a42ddc160a713bf0efa Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Mon, 26 Jan 2026 08:02:47 -0600 Subject: [PATCH 04/19] update hermes --- ditto_ms.py | 653 +++++++++++++++++++++++++++-------------------- environment.yml | 8 + hermes.py | 351 +++++++++++++++++++++++++ inc/data.py | 14 +- inc/diffus.py | 100 ++++++-- inc/test_ms.py | 8 +- requirements.txt | 65 +++++ 7 files changed, 877 insertions(+), 322 deletions(-) create mode 100644 environment.yml create mode 100644 hermes.py create mode 100644 requirements.txt diff --git a/ditto_ms.py b/ditto_ms.py index 055b3b5..9ec7c1b 100644 --- a/ditto_ms.py +++ b/ditto_ms.py @@ -150,6 +150,226 @@ def lik(self, Y): # Y: (T+1, nodes, samples) ) ).sum(dim = 0) # (samples,) return lik, zI0, zR0, zI, zR # (samples,) + + def _lik_seg_local(self, Y, zI, zR, TL, TR): + """ + Segment log-prob under the *original DITTO local-support* backward sampler. + This is essentially the original `lik()` but restricted to t in [TL, TR-1]. + + Parameters + ---------- + Y : LongTensor, (T+1, nodes, samples) + zI, zR : FloatTensor, (T, nodes, samples) (already clamped / combined) + TL, TR : int segment endpoints (TL < TR) + """ + n_samples = Y.size(dim=2) + L = int(TR - TL) + if L <= 0: + return torch.zeros(n_samples, dtype=torch.float32, device=self.device) + + # Slice the segment (TL..TR) + Yseg = Y[TL: TR + 1] # (L+1, nodes, samples) + zIseg = zI[TL: TR] # (L, nodes, samples) + zRseg = zR[TL: TR] # (L, nodes, samples) + + # ------------------------- + # R -> I (backward) + # ------------------------- + qR = torch.sigmoid(zRseg) # (L, nodes, samples) + lR1 = torch_log(qR) + lR0 = torch_log(1.0 - qR) + with torch.no_grad(): + mskR = (Yseg[1:] == SIR_STATES.R) # (L, nodes, samples) + trsR = (Yseg[:-1] != SIR_STATES.R) # (L, nodes, samples) + + # ------------------------- + # I -> S (backward) with DITTO ordering trick + # ------------------------- + zI_, uid = zIseg.sort(dim=1, descending=True) # (L, nodes, samples) + qI = torch.sigmoid(zI_) # (L, nodes, samples) + lI1 = torch_log(qI) + lI0 = torch_log(1.0 - qI) + + with torch.no_grad(): + # NOTE: this follows your existing `lik()` implementation style. + mskI = ((Yseg[1:] >= SIR_STATES.I) & (Yseg[:-1] <= SIR_STATES.I)).flatten() + trsI = ((Yseg[:-1] != SIR_STATES.I)).flatten() + + rem = torch.where(mskI, self.rem.expand(L, -1, n_samples).flatten(), self.n_inf) + ptr = torch.arange(L, dtype=torch.long, device=self.device).unsqueeze(dim=1) * self.n_nodes # (L,1) + + for i in range(uid.size(dim=1)): + uidi = (ptr + uid[:, i]).flatten() * n_samples # (L*samples) + mski = mskI[uidi] + if mski.max(): + trsi = trsI[uidi] + + vids, degi = [], [0] + for t in range(L): + for j in range(n_samples): + u = uid[t, i, j] + vids.append((t * self.n_nodes + u.unsqueeze(dim=0)) * n_samples) + vid = self.neighbs[u.item()] + vids.append((t * self.n_nodes + vid) * n_samples) + degi.append(vid.size(dim=0) + 1) + + vids = torch.cat(vids, dim=0) + degi = torch.tensor(degi, dtype=torch.long, device=self.device) + indptr = degi.cumsum(dim=0) + degi = degi[1:] + + rems = rem.flatten()[vids] + opti = (pysc.segment_min_csr(src=rems, indptr=indptr)[0] > 1) + + rem.flatten()[vids] = torch.where( + mski.repeat_interleave(repeats=degi), + torch.where(trsi.repeat_interleave(repeats=degi), rems - 1, self.n_inf), + rems, + ) + mskI[uidi] &= opti + + lik = ( + torch.where(mskR, torch.where(trsR, lR1, lR0), self.zero).reshape(-1, n_samples) + + torch.where( + mskI.reshape(-1, n_samples), + torch.where( + trsI.reshape(-1, n_samples), + lI1.reshape(-1, n_samples), + lI0.reshape(-1, n_samples), + ), + self.zero, + ) + ).sum(dim=0) + + return lik + + def _lik_seg_a1(self, Y, zI, zR, TL, TR, eps=1e-6): + S, I, R = SIR_STATES.S, SIR_STATES.I, SIR_STATES.R + device = self.device + n_samples = Y.shape[2] + p_seg = self.zero.expand(n_samples).clone() + + + yL = Y[TL] + blocked = (yL == R) + src = (yL == I) + + # IMPORTANT: only t in (TL, TR) i.e. TL+1 ... TR-1 (consistent with _samp_seg) + for t in range(TL + 1, TR): + d = t - TL + y_t = Y[t] + y_next = Y[t + 1] + + qR = torch.sigmoid(zR[t]).clamp(eps, 1.0 - eps) # (n_nodes, 1) + p_add = (1.0 - torch.sigmoid(zI[t])).clamp(eps, 1.0 - eps) # (n_nodes, 1) + + pool = (y_next != S) & (~blocked) + A_true = (y_t != S) & pool + + # replay growth to compute log prob of A_true + A = src & pool + frontier = A.clone() + decided = A.clone() + + logp_A = torch.zeros(y_t.shape[1], device=device) + + for _ in range(d): + cand = (torch.sparse.mm(self.adj, frontier.float()) > 0) & pool & (~decided) + if not cand.any(): + break + + add = cand & A_true + logp_A += (add.float() * torch_log(p_add) + (cand & ~add).float() * torch_log(1 - p_add)).sum(0) + + A |= add + frontier = add + decided |= cand + + # if A_true contains nodes never reached by growth => prob 0 + missing = A_true & (~A) + if missing.any(): + # set those samples to -inf + bad = missing.any(dim=0) + logp_A[bad] = -float("inf") + + # candR likelihood + candR = A_true & (y_next == R) + isI = (y_t == I) + toI = candR & isI + toR = candR & (~isI) # should be R + + logp_R = (toI.float() * torch_log(qR) + toR.float() * torch_log(1 - qR)).sum(0) + + # if y_next==I and in A_true, must be I + badI = (A_true & (y_next == I) & (y_t != I)).any(dim=0) + if badI.any(): + logp_R[badI] = -float("inf") + + p_seg += (logp_A + logp_R) + + return p_seg + + def lik_ms(self, Y, obs_time): + """ + Multi-snapshot proposal likelihood matching `samp_ms()`. + """ + assert Y.size(dim=0) == self.T + 1, "lik_ms expects Y with shape (T+1, nodes, samples)" + n_samples = Y.size(dim=2) + + # sanitize obs_time: unique, within (0..T], and must include T + obs_time = sorted({int(t) for t in obs_time if 0 < int(t) <= self.T}) + if self.T not in obs_time: + obs_time.append(self.T) + K = len(obs_time) + + # Condition on all observed snapshots: concat along the "samples" dimension + y_cond = torch.cat([Y[t] for t in obs_time], dim=1) + + # Forward once for all conditioning blocks + zI0, zR0, zI, zR = self.forward(y_cond, orig=True) # (T, nodes, samples*K) + + # reshape to (T, nodes, K, samples) + zI0 = zI0.contiguous().view(self.T, self.n_nodes, K, n_samples) + zR0 = zR0.contiguous().view(self.T, self.n_nodes, K, n_samples) + zI = zI.contiguous().view(self.T, self.n_nodes, K, n_samples) + zR = zR.contiguous().view(self.T, self.n_nodes, K, n_samples) + + # detach+clamp leaf trick (same training mechanism as original lik()) + zI = zI.clone().detach().requires_grad_(True) + zI.retain_grad() + zR = zR.clone().detach().requires_grad_(True) + zR.retain_grad() + + lik = torch.zeros(n_samples, dtype=torch.float32, device=self.device) + + segL = [0] + obs_time[:-1] + segR = obs_time + + for i in range(K): + TL, TR = segL[i], segR[i] + + # logits conditioned on right endpoint snapshot y_TR (index i) + zI_R = zI[:, :, i, :] # (T, nodes, samples) + zR_R = zR[:, :, i, :] + + # combine left+right logits for TL>0 as in samp_ms() + if (TL > 0) and (K > 1): + zI_L = zI[:, :, i - 1, :] + zR_L = zR[:, :, i - 1, :] + zI_seg = self.clamp_z(zI_R + zI_L) + zR_seg = self.clamp_z(zR_R + zR_L) + else: + zI_seg = zI_R + zR_seg = zR_R + + if TL == 0: + lik = lik + self._lik_seg_local(Y, zI_seg, zR_seg, TL=TL, TR=TR) + else: + # _lik_seg_a1() only returns `lik` and does not accept return_ok / invalid_to_neg_inf + lik = lik + self._lik_seg_a1(Y, zI_seg, zR_seg, TL=TL, TR=TR) + + return lik, zI0, zR0, zI, zR + @torch.no_grad() def clamp_grad(self, z0, grad): return torch.where(z0 < self.zlim, torch.where(z0 > -self.zlim, grad, F.relu(grad)), -F.relu(-grad)) @@ -296,301 +516,153 @@ def _samp_step(self, y, zI, uid, zR, t, compute_lik=False, yL=None): else: return y - @torch.no_grad() - def _samp_seg(self, yR, zI, uid, zR, n_samples, TL, TR, yL=None, compute_lik=False): - """ - Sample ONE segment (TL, TR] in *reverse* temporal order (backward sampling). - - Segment definition: - - Left endpoint time TL (may be observed / clamped if yL is provided) - - Right endpoint time TR (always observed here; yR is the snapshot at TR) - - We generate snapshots for times: TR-1, TR-2, ..., TL (if yL is None) or TL+1 (if yL is fixed) - - Why segment-wise? - In multi-snapshot DITTO, we must satisfy *all* observed snapshots exactly. - We sample each segment backward from the fixed right endpoint y_TR. However, in the multi-snapshot - setting we also need a *left feasibility* guarantee: sampled states must still be extendable to - match the left observed snapshot y_TL. This is the left-extendability hard constraint Ext(t). - - Parameters - ---------- - yR : LongTensor, shape (nodes,) - Fixed right endpoint snapshot y_{TR}. - zI : FloatTensor, shape (T, nodes, 1) - Proposal logits (already sorted along nodes dim) controlling backward I->S decisions. - NOTE: must be consistent with uid. - uid : LongTensor, shape (T, nodes) - Node indices sorted by descending zI at each time t (DITTO's ordering trick). - zR : FloatTensor, shape (T, nodes, 1) - Proposal logits controlling backward R->I decisions (NOT sorted; original node order). - n_samples : int - How many independent histories to sample in parallel. - TL, TR : int - Segment endpoints (TL < TR). - yL : LongTensor or None - If provided, this is the *observed* left endpoint snapshot y_{TL} (hard constraint). - In that case we will NOT sample time TL; we clamp it to yL. - compute_lik : bool - If True, also return log-probability under the proposal (needed by M-H acceptance). - - Returns - ------- - Y_seg : LongTensor, shape (TR-TL, nodes, samples) - The sampled segment snapshots in *forward* time indexing within the segment: - Y_seg[k] corresponds to time (TL + k), for k=0..(TR-TL-1). - If yL is provided, then Y_seg[0] == yL is included (clamped). - lik_seg : FloatTensor, shape (samples,) (only if compute_lik=True) - Sum of log-probabilities of all backward steps performed inside this segment. - """ - L = TR - TL # number of time indices in [TL, TR) that we store in Y_seg - Y = torch.empty(L, self.n_nodes, n_samples, dtype=torch.long, device=self.device) - - if compute_lik: - # segment log-likelihood under the proposal Q_theta - lik = torch.zeros(n_samples, dtype=torch.float, device=self.device) - - # Current "right" snapshot y_{t+1}. Start from fixed right endpoint y_{TR}. - y = yR.unsqueeze(dim=1).expand(-1, n_samples) # (nodes, samples) - - # --------------------------------------------------------------------- - # Precompute yL-dependent tensors once per segment. - # --------------------------------------------------------------------- - if yL is not None: - yL_col = yL.unsqueeze(dim=1) # (nodes, 1) for monotonicity check x >= y_TL - yL_not_R = (yL != SIR_STATES.R).unsqueeze(dim=1) # (nodes, 1) exclude nodes fixed to R at TL - src = (yL == SIR_STATES.I).unsqueeze(dim=1) # (nodes, 1) infection sources at TL - - def ext_ok_fast(x, t): - """ - Faster ext_ok that reuses yL-related precomputations. - x: (nodes, samples) candidate snapshot at time t - """ - d = t - TL - ok = (x >= yL_col).all(dim=0) # (samples,) - if not ok.max(): - return ok - A = (x != SIR_STATES.S) & yL_not_R # (nodes, samples) - reach = src.expand(-1, x.size(dim=1)) & A # (nodes, samples) - for _ in range(d): - nbr = (torch.sparse.mm(self.adj, reach.float()) > 0) & A - new_reach = reach | nbr - if torch.equal(new_reach, reach): - reach = new_reach - break - reach = new_reach - Iset = (x == SIR_STATES.I) & yL_not_R - Rset = (x == SIR_STATES.R) & yL_not_R - ok = ok & ~(Iset & ~reach).any(dim=0) - ok = ok & ~(Rset & ~reach).any(dim=0) - return ok - else: - ext_ok_fast = None - - # --------------------------------------------------------------------- - # Segment-level empty-support check: the observed right endpoint itself - # must be extendable from the observed left endpoint (when present). - # --------------------------------------------------------------------- - if ext_ok_fast is not None: - okR = ext_ok_fast(yR.unsqueeze(dim=1), TR) # (1,) - if not bool(okR.item()): - n_src = int((yL == SIR_STATES.I).sum().item()) - n_fixR = int((yL == SIR_STATES.R).sum().item()) - n_yR_I = int((yR == SIR_STATES.I).sum().item()) - n_yR_R = int((yR == SIR_STATES.R).sum().item()) - mono = bool((yR >= yL).all().item()) - raise RuntimeError( - "[DITTO-MS] Segment infeasible (empty support) under hard constraints. " - f"Segment (TL={TL}, TR={TR}, len={TR-TL}). " - f"Monotonic(y_TR>=y_TL)={mono}. " - f"#src_I@TL={n_src}, #fixed_R@TL={n_fixR}, #I@TR={n_yR_I}, #R@TR={n_yR_R}. " - "This typically means the observations cannot be bridged on the graph within the time budget " - "(e.g., src is empty/small, isolated targets, or timestamps over-constrain the diffusion)." - ) - - # --------------------------------------------------------------------- - # First segment (TL=0) has no left-extendability constraint; keep original sampler. - # For subsequent segments, use constructive hop-layer growth to enforce Ext(t) by construction. - # --------------------------------------------------------------------- - if yL is None: - # no left constraint, sample purely by original backward local-support steps - for k in range(L - 1, -1, -1): - t = TL + k + def _samp_seg(self, yL, yR, zI_sorted, uid, zR, TL, TR, compute_lik=False): + S, I, R = self.S, self.I, self.R + n_samples = yR.shape[1] + device = self.device + + Y = torch.empty(TR - TL + 1, self.n_nodes, n_samples, dtype=torch.long, device=device) + lik = 0.0 if compute_lik else None + + # Clamp left endpoint + yL = yL.unsqueeze(dim=1).expand(-1, n_samples) + yR = yR.unsqueeze(dim=1).expand(-1, n_samples) + Y[0] = yL + Y[TR - TL] = yR + + # TL=0 : keep original DITTO step + if TL == 0: + y = yR + for t in range(TR - 1, 0, -1): + y, p_x = self._samp_step(y, zI_sorted[t], uid[t], zR[t], compute_lik) + Y[t] = y if compute_lik: - y, lik_t = self._samp_step(y, zI, uid, zR, t, compute_lik=True, yL=None) - lik = lik + lik_t - else: - y = self._samp_step(y, zI, uid, zR, t, compute_lik=False, yL=None) - Y[k] = y - if compute_lik: - return Y.detach().clone(), lik.detach().clone() - else: - return Y.detach().clone() - - # --------------------------------------------------------------------- - # Constructive extendable sampler (Scheme A1): - # At each time t in (TL, TR): - # - pool := nodes with y_{t+1} != S and yL != R - # - build A_t (non-S set) by hop layers from src within pool: - # include all nodes within distance <= d-1 - # optionally include some nodes at exact distance d - # force-include distance-d nodes that are needed to infect distance-(d+1) nodes - # - set y_t outside A_t to S (or R if blocked) - # - inside A_t, decide I/R using zR, but force I on nodes needed as infection sources - # This eliminates the rejection loop for Ext(t). - # --------------------------------------------------------------------- - - # Build unsorted zI (needed to derive p_add on boundary layer in original node order). - # zI is sorted along nodes dim with permutation uid. - zI_sorted_ = zI.squeeze(dim=2) # (T, nodes) - zI_unsorted = torch.empty_like(zI_sorted_) # (T, nodes) - for tt in range(self.T): - zI_unsorted[tt, uid[tt]] = zI_sorted_[tt] - - # Precompute fixed masks from left observation. - blocked = (yL == SIR_STATES.R).unsqueeze(dim=1) # (nodes, 1) - yL_not_R = ~blocked - src = (yL == SIR_STATES.I).unsqueeze(dim=1) # (nodes, 1) - - # We do NOT sample time TL itself; clamp it to yL. - k_min = 1 - - for k in range(L - 1, k_min - 1, -1): - t = TL + k - d = t - TL # hop budget for Ext(t) + lik += p_x + return Y, lik - y_next = y # y_{t+1}, shape (nodes, samples) + # ------------------------- + # TL > 0 : Route-B sampler + # ------------------------- - # pool = {u: y_{t+1,u} != S} \ blocked - pool = (y_next != SIR_STATES.S) & yL_not_R # (nodes, samples) + # unsort zI for p_add lookup (keep your original logic) + zI_unsorted = torch.empty_like(zI_sorted).squeeze(dim=2) # (T, n_nodes) + for tt in range(self.T): + zI_unsorted[tt].scatter_(dim=0, index=uid[tt], src=zI_sorted[tt].squeeze(dim=1)) - # If any sample has non-empty pool but no src in pool, segment is infeasible for that sample. - src_in_pool = (src.expand(-1, n_samples) & pool).any(dim=0) # (samples,) - if (~src_in_pool & pool.any(dim=0)).any(): - raise RuntimeError( - "[DITTO-MS] Constructive sampler hit infeasible intermediate state: " - f"at time t={t} (TL={TL},TR={TR}), pool non-empty but src not in pool for some samples. " - "This indicates either inconsistent observations or a bug in monotonic clamping." - ) + blocked = (yL == R) + src = (yL == I) - # ----------------------------------------------------------------- - # Hop-layer BFS within pool from src, up to d+1 layers. - # We need: - # - reach_{d-1}: nodes within dist <= d-1 (must be non-S at time t to support outer growth) - # - layer_d: nodes at exact dist d (optional, but some are forced) - # - layer_{d+1}: nodes at exact dist d+1 (cannot be non-S at time t; must be new at t+1) - # ----------------------------------------------------------------- - frontier = src.expand(-1, n_samples) & pool # layer 0 - reach = frontier.clone() - reach_dminus1 = frontier.clone() # will be overwritten if d-1 >= 1 - layer_d = torch.zeros_like(frontier) - layer_d1 = torch.zeros_like(frontier) - - # Special: if d-1 == 0, then reach_dminus1 is just layer0. - # We'll record reach after step (d-1) as reach_dminus1. - for h in range(1, d + 2): # compute layers 1..d+1 - nbr = (torch.sparse.mm(self.adj, frontier.float()) > 0) & pool & (~reach) - frontier = nbr - reach = reach | frontier - if h == d - 1: - reach_dminus1 = reach.clone() - if h == d: - layer_d = frontier.clone() - if h == d + 1: - layer_d1 = frontier.clone() - - # If d == 1, loop sets reach_dminus1 when h==0 not visited; keep as layer0. - if d == 1: - reach_dminus1 = src.expand(-1, n_samples) & pool - - # Nodes at dist <= d-1 are always included in A_t (conservative constructive core). - A = reach_dminus1.clone() - - # Force-include distance-d nodes that are adjacent to distance-(d+1) nodes, - # because those layer_{d+1} nodes must be infected at t+1 and need an I neighbor at time t. - if d >= 1: - bnd_need = layer_d & (torch.sparse.mm(self.adj, layer_d1.float()) > 0) - else: - bnd_need = torch.zeros_like(layer_d) + for t in range(TR - 1, TL, -1): + d = t - TL + y_next = Y[t - TL + 1] # (n_nodes, n_samples) - # Optional boundary nodes at dist d that are NOT needed for layer_{d+1}. - bnd_opt = layer_d & (~bnd_need) + # probabilities + qR = torch.sigmoid(zR[t]) # (n_nodes, 1) + p_add = 1.0 - torch.sigmoid(zI_unsorted[t]).unsqueeze(1) # (n_nodes, 1) + # (optional) clamp to avoid exactly 0/1 + qR = qR.clamp(self.eps, 1.0 - self.eps) + p_add = p_add.clamp(self.eps, 1.0 - self.eps) - # Sample add decision for optional boundary nodes using p_add = 1 - sigmoid(zI_unsorted[t]). - # (High qI => more likely to become new, so lower include prob.) - p_add = (1.0 - torch.sigmoid(zI_unsorted[t]).unsqueeze(dim=1)).clamp(1e-6, 1.0 - 1e-6) # (nodes,1) - if bnd_opt.any(): - rnd = torch.rand(self.n_nodes, n_samples, device=self.device) - bnd_add = bnd_opt & (rnd <= p_add.expand(-1, n_samples)) - else: - bnd_add = torch.zeros_like(bnd_opt) - - # Final A_t - A = A | bnd_need | bnd_add - - # ----------------------------------------------------------------- - # Given A_t, construct y_t: - # - blocked nodes (yL==R): force R - # - nodes not in A and not blocked: S - # - nodes in A: - # if y_{t+1}==I => force I - # if y_{t+1}==R => sample I/R with qR, but force I if needed as infection source - # ----------------------------------------------------------------- - y_t = torch.full((self.n_nodes, n_samples), SIR_STATES.S, dtype=torch.long, device=self.device) - y_t = torch.where(blocked.expand(-1, n_samples), SIR_STATES.R, y_t) - - # new nodes are those in pool but not in A (they are S at time t, become non-S at t+1) - new = pool & (~A) + pool = (y_next != S) & (~blocked) - # Any node in A adjacent to any new node must be infected at time t to enable infection. - need_source = A & (torch.sparse.mm(self.adj, new.float()) > 0) + # ------------------------------------------------------------ + # Step 1: A_t generation WITHOUT forcing reach_{<=d-1} + # (outward growth from src for d hops) + # ------------------------------------------------------------ + A = src & pool + frontier = A.clone() + decided = A.clone() # nodes whose add/not-add decision has been made - # Force I where y_{t+1}==I - force_I_from_next = A & (y_next == SIR_STATES.I) + if compute_lik: + logp_A = torch.zeros(n_samples, device=device) - # Candidate nodes with y_{t+1}==R that are not forced infection sources. - candR = A & (y_next == SIR_STATES.R) & (~need_source) & (~force_I_from_next) + for _ in range(d): + cand = (torch.sparse.mm(self.adj, frontier.float()) > 0) & pool & (~decided) + if not cand.any(): + break - # Sample I/R for candR using qR (prob of R->I backward => I at time t). - if candR.any(): - qR_t = torch.sigmoid(zR[t]).clamp(1e-6, 1.0 - 1e-6) # (nodes,1) - rndR = torch.rand(self.n_nodes, n_samples, device=self.device) - isI = candR & (rndR <= qR_t.expand(-1, n_samples)) - # Set sampled states - y_t = torch.where(isI, SIR_STATES.I, y_t) - y_t = torch.where(candR & (~isI), SIR_STATES.R, y_t) - if compute_lik: - logq = torch_log(qR_t).expand(-1, n_samples) - log1q = torch_log(1.0 - qR_t).expand(-1, n_samples) - lik = lik + (torch.where(isI, logq, self.zero) + torch.where(candR & (~isI), log1q, self.zero)).sum(dim=0) - else: + u = torch.rand(self.n_nodes, n_samples, device=device) + add = cand & (u <= p_add) # Bernoulli(p_add) if compute_lik: - pass + logp_A += (add.float() * torch_log(p_add) + (cand & ~add).float() * torch_log(1 - p_add)).sum(0) + + A |= add + frontier = add + decided |= cand + + # ------------------------------------------------------------ + # Step 2: sample states in A, but DO NOT force all neighbors of new + # Enforce: each new has >=1 infected neighbor (only if otherwise 0-prob) + # ------------------------------------------------------------ + x_t = torch.full_like(y_next, S) + x_t[blocked] = R + + # forced I if y_next==I and in A + mI = A & (y_next == I) + x_t[mI] = I + + # candR nodes: y_next==R and in A => sample I/R via qR + candR = A & (y_next == R) + uR = torch.rand(self.n_nodes, n_samples, device=device) + toI = candR & (uR <= qR) + toR = candR & (~toI) + x_t[toI] = I + x_t[toR] = R - # Force I for infection sources and nodes that are I at t+1 - y_t = torch.where(need_source | force_I_from_next, SIR_STATES.I, y_t) + if compute_lik: + logp_R = (toI.float() * torch_log(qR) + toR.float() * torch_log(1 - qR)).sum(0) - # ----------------------------------------------------------------- - # Proposal likelihood contributions from boundary add decisions. - # We only account for sampled optional boundary nodes (bnd_opt). - # Forced inclusions (reach<=d-1 and bnd_need) are deterministic (log prob 0). - # ----------------------------------------------------------------- - if compute_lik and bnd_opt.any(): - p = p_add.expand(-1, n_samples) - logp = torch_log(p) - log1p = torch_log(1.0 - p) - lik = lik + (torch.where(bnd_add, logp, self.zero) + torch.where(bnd_opt & (~bnd_add), log1p, self.zero)).sum(dim=0) + # new nodes + new = pool & (~A) - # Commit and step left - Y[k] = y_t - y = y_t + # ---- enforce infection-source constraint minimally ---- + # For each sample: if a new node has no infected neighbor, pick ONE neighbor in A and flip to I + # (only when otherwise impossible / zero-prob forward) + I_mask = (x_t == I) + neighI = (torch.sparse.mm(self.adj, I_mask.float()) > 0) + bad_new = new & (~neighI) # nodes that violate ">=1 infected neighbor" + if bad_new.any(): + # candidates that could be flipped to I: in A, and (y_next==I already I) OR (y_next==R and in candR) + # (y_next==I in A are already I; so only need consider A & (y_next==R) that are currently R) + fixable = A & (y_next == R) + + # For each bad_new node u, choose one neighbor v from fixable ∩ N(u) to flip to I + # If none exists, that sample is infeasible under model (posterior prob 0), so leave as is (will be rejected later if you have a checker) + neigh_fixable = (torch.sparse.mm(self.adj, fixable.float()) > 0) + # We flip per bad_new node by sampling one neighbor index. + # Implementation trick: do one pass "greedy-random" by picking first available neighbor per (u,sample) + # (this avoids heavy per-node loops and still gives each choice positive prob if you randomize tie-breaking) + # Here: random tie-breaking by multiplying adjacency mask with random noise. + adj_dense = self.adj.to_dense() # (n_nodes, n_nodes) maybe too big; if too big, replace with sparse gather kernels + # NOTE: if graph is large, do not materialize dense. In that case implement sparse neighbor sampling separately. + + # For minimal code change, keep dense only if feasible in your scale. + noise = torch.rand(self.n_nodes, self.n_nodes, device=device) + # candidates matrix: (u,v,sample) => u in bad_new, v in fixable neighbor + # build neighbor mask (u,v) then apply for each sample + nb_mask = (adj_dense > 0) + + for s in range(n_samples): + bad_u = bad_new[:, s].nonzero(as_tuple=False).flatten() + if bad_u.numel() == 0: + continue + fix_v = fixable[:, s] + # for each bad u, pick v maximizing noise among allowed neighbors + for u_node in bad_u.tolist(): + allowed = nb_mask[u_node] & fix_v + if allowed.any(): + v_idx = (noise[u_node] * allowed.float()).argmax().item() + x_t[v_idx, s] = I + + # optional: if you want strict soundness, you can assert-check again and resample/reject, + # but for minimal code change we just repair as above. + + Y[t - TL] = x_t - # Clamp left endpoint snapshot - Y[0] = yL.unsqueeze(dim=1).expand(-1, n_samples) + if compute_lik: + lik += (logp_A + logp_R) - if compute_lik: - return Y.detach().clone(), lik.detach().clone() - else: - return Y.detach().clone() + return Y, lik @torch.no_grad() def samp_ms(self, y, zI, zR, n_samples, obs_time, compute_lik=False): @@ -671,22 +743,37 @@ def _pick(z, i): return Y.detach().clone() -def q_loss(q_net, data, I0, bpar, n_samples): +def q_loss(q_net, data, I0, bpar, n_samples, obs_time): T = data.T.item() n_nodes = data.num_nodes - Y = diffus_gen(T = T, n_nodes = n_nodes, edge_index = data.edge_index, I0 = I0, n_samples = n_samples, pI = bpar.pI, pR = bpar.pR) # (T+1, nodes, samples) - q_liks, zI0, zR0, zI, zR = q_net.lik(Y = Y) # (samples,) + Y = diffus_gen( + T=T, + n_nodes=n_nodes, + edge_index=data.edge_index, + I0=I0, + n_samples=n_samples, + pI=bpar.pI, + pR=bpar.pR, + ) # (T+1, nodes, samples) + + q_liks, zI0, zR0, zI, zR = q_net.lik_ms(Y=Y, obs_time=obs_time) return -q_liks.mean(), zI0, zR0, zI, zR def q_train(data, bpar, args): I0 = (data.y[:, 0] == 1).long().sum().item() + + # Keep training obs_time consistent with main()/t_mcmc. + obs_time = [int(t) for t in args.obs_time.split(',') if t] + obs_time.append(data.T.item()) + obs_time = sorted({t for t in obs_time if 0 < t <= data.T.item()}) + q_net = QNet.make(data, args) q_net.train() - opt = optim.AdamW(q_net.parameters(), lr = args.q_lr) + opt = optim.AdamW(q_net.parameters(), lr=args.q_lr) pbar = trange(1, args.q_steps + 1) for step in pbar: opt.zero_grad() - loss, zI0, zR0, zI, zR = q_loss(q_net, data, I0, bpar, args.q_samples) + loss, zI0, zR0, zI, zR = q_loss(q_net, data, I0, bpar, args.q_samples, obs_time=obs_time) pbar.set_description(f'[step={step}] loss={loss.item():.4f}') q_net.backward(loss, zI0, zR0, zI, zR) opt.step() diff --git a/environment.yml b/environment.yml new file mode 100644 index 0000000..cc1e9af --- /dev/null +++ b/environment.yml @@ -0,0 +1,8 @@ +name: ditto-gpu +channels: + - defaults +dependencies: + - python=3.11.14 + - pip + - git +prefix: C:\Users\z1585\anaconda3\envs\ditto-gpu diff --git a/hermes.py b/hermes.py new file mode 100644 index 0000000..bb99f34 --- /dev/null +++ b/hermes.py @@ -0,0 +1,351 @@ +import sys + +from inc.diffus import * +from inc.nn import * +from inc.test import * + +def get_args(): + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type = str, help = 'dataset name') + parser.add_argument('--seed', type = int, help = 'random seed') + parser.add_argument('--data_dir', type = str, help = 'dataset folder') + parser.add_argument('--output', type = str, help = 'output file name') + parser.add_argument('--device', type = torch.device, help = 'torch device') + parser.add_argument('--b_pI0', type = float, help = 'initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type = float, help = 'initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type = int, help = 'optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type = float, help = 'learning rate in diffusion parameter estimation') + parser.add_argument('--q_steps', type = int, help = 'training steps for the proposal model') + parser.add_argument('--q_lr', type = float, help = 'learning rate for the proposal model') + parser.add_argument('--q_hid', type = int, help = 'hidden size of the proposal model') + parser.add_argument('--q_gnn', type = int, help = 'number of layers of the GNN in the proposal model') + parser.add_argument('--q_mlp', type = int, help = 'number of layers of the MLP in the proposal model') + parser.add_argument('--q_samples', type = int, help = 'sample size to estimate the loss function of the proposal model') + parser.add_argument('--q_zlim', type = int, help = 'a hyperparameter to stablize gradient') + parser.add_argument('--p_coef', type = float, help = 'the coefficient gamma in the initial distribution P[y_0]') + parser.add_argument('--t_samples', type = int, help = 'MCMC sample size') + parser.add_argument('--t_steps', type = int, help = 'MCMC steps') + parser.add_argument('--t_keep', type = float, help = 'moving average in MCMC') + parser.add_argument('--obs_time', type = str, default = '', help = 'extra observed snapshot times, comma-separated, e.g., 5,7,9') + args = parser.parse_args() + return args + +class QNet(nn.Module): + @classmethod + def make(cls, data, obs_time, args): + return cls( + eidx = data.edge_index, + T = data.T.item(), + n_obs = len(obs_time), + hid = args.q_hid, + gnn = args.q_gnn, + mlp = args.q_mlp, + n_nodes = data.num_nodes, + zlim = args.q_zlim, + ).to(args.device) + def __init__(self, eidx, T, n_obs, hid, gnn, mlp, n_nodes, zlim): + super().__init__() + self.eidx = eidx + self.device = self.eidx.device + self.n_nodes = n_nodes + self.n_inf = self.n_nodes + 2 + self.n_edges = self.eidx.size(dim = 1) + self.zlim = zlim + # Maximum rejection rounds per backward step in multi-snapshot segment sampling. + # This avoids infinite loops when a segment has empty/tiny support under hard constraints. + self.T = T + self.hid = int(hid) + self.gnn_dep = int(gnn) + self.mlp_dep = int(mlp) + self.w = nn.Parameter(data = torch.randn((self.n_edges, self.hid), dtype = torch.float32, device = self.device), requires_grad = True) + self.gnn = GNN(v_in = n_obs, e_in = self.hid, hid = self.hid, dep = self.gnn_dep) + self.mlp = MLP([self.hid] * self.mlp_dep + [2 * self.T]) + self.rem = (pyg.utils.degree(self.eidx[1], num_nodes = self.n_nodes).long().unsqueeze(dim = 1) + 1).detach().clone() # (nodes, 1) + self.neighbs = [[] for u in range(self.n_nodes)] + for i in range(self.n_edges): + self.neighbs[self.eidx[0, i].item()].append(self.eidx[1, i].item()) + for u in range(self.n_nodes): + self.neighbs[u] = torch.tensor(self.neighbs[u], dtype = torch.long, device = self.device) + self.adj = torch.sparse_coo_tensor( + indices = torch.stack([self.eidx[1], self.eidx[0]], dim = 0), + values = torch.ones(self.n_edges, dtype = torch.float, device = self.device), + size = (self.n_nodes, self.n_nodes), + ).coalesce() + self.zero = torch.tensor(0., dtype = torch.float, device = self.device) + def clamp_z(self, z): + return z.clamp(-self.zlim, self.zlim) + def forward(self, y, orig = False): # y: (nodes, samples, obs) + n_nodes, n_samples, n_obs = y.size() + y = y.permute(1, 0, 2).flatten(end_dim = 1) # (samples*nodes, obs) + eidx = (self.eidx.unsqueeze(dim = 1) + n_nodes * torch.arange(n_samples, dtype = torch.long, device = y.device).unsqueeze(dim = -1)).reshape((2, -1)) # (2, samples*edges) + w = self.w.repeat(n_samples, 1) # (samples, hid) + z, e = self.gnn(y.float(), eidx, w) + z = self.mlp(z) # (samples*nodes, 2*T) + z = z.T.reshape((2 * self.T, n_samples, -1)) # (2*T, samples, nodes) + zI, zR = z[: self.T], z[self.T :] # (T, samples, nodes) + zI, zR = zI.transpose(1, 2), zR.transpose(1, 2) # (T, nodes, samples) + if orig: + return zI, zR, self.clamp_z(zI), self.clamp_z(zR) + else: + return self.clamp_z(zI), self.clamp_z(zR) + def _lik_step(self, y0, y1, lI1, lI0, lR1, lR0, yL=None, reach=None): # y*, l*, reach: (nodes, sampls) + n_nodes, n_samples = y0.shape + lik = self.zero + # unreachable has log1=0 + # R->I + msk = (y1 == SIR_STATES.R) # (nodes, samples) + if yL is not None: + msk = msk & (yL != SIR_STATES.R) & reach + lik = lik + torch.where(msk, torch.where(y0 != SIR_STATES.R, lR1, lR0), self.zero) + # I->S + uid = lI1.argsort(dim = 0, descending = True) # (nodes, samples) + msk = (y1 == SIR_STATES.I) | (msk & (y0 != SIR_STATES.R)) # (nodes, samples) + rem = torch.where(msk, (reach.long() if reach is not None else 1) + torch.sparse.mm(self.adj, msk.float()).long(), self.n_inf) # (nodes, samples) + cols = torch.arange(n_samples, dtype = torch.long, device = rem.device) + for i, u in enumerate(uid): + rem_v = torch.full((self.n_nodes, n_samples), self.n_inf, dtype=rem.dtype, device=rem.device) + rem_v = rem_v.scatter_reduce(dim=0, index=self.eidx[0, :, None].expand(-1, n_samples), + src=rem[self.eidx[1]], reduce="amin", include_self=True) + rem_v = rem_v.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) # (samples,) + rem_u = rem.gather(dim = 0, index = u.unsqueeze(dim=0)).squeeze(dim=0) # (samples,) + opt = (rem_u > 1) & (rem_v > 1) # (samples,) + if yL is not None: + opt = opt & (yL.gather(dim = 0, index = u.unsqueeze(dim=0)).squeeze(dim=0) != SIR_STATES.I) + msk_u = msk.gather(dim = 0, index = u.unsqueeze(dim=0)).squeeze(dim=0) # (samples,) + msk_opt = msk_u & opt # (samples,) + lik = lik + torch.where(msk_opt, torch.where(y0 == SIR_STATES.S, lI1, lI0), self.zero) + trs = (y0.gather(dim = 0, index = u.unsqueeze(dim=0)).squeeze(dim=0) == SIR_STATES.S) # (samples,) + rem = rem - torch.zeros_like(rem).index_put( + (self.eidx[1, :, None].expand(-1, n_samples).flatten(), cols[None].expand(self.n_edges, -1).flatten()), + ((msk_u & trs) & (self.eidx[0, :, None] == u)).flatten().long(), # (edges * samples,) + accumulate = False) + rem = rem.index_put((u, cols), torch.where(msk_u, torch.where(trs, rem_u - 1, self.n_inf), rem.gather(dim = 0, index = u.unsqueeze(dim=0)).squeeze(dim=0))) + msk = msk.index_put((u, cols), msk_opt) + return lik + def lik_ms(self, Y, obs_time): #(T+1, nodes, samples) + assert Y.size(dim=0) == self.T + 1, "lik_ms expects Y with shape (T+1, nodes, samples)" + n_samples = Y.size(dim=2) + K = len(obs_time) + y_cond = torch.stack([Y[t] for t in obs_time], dim=2) # (nodes, samples, obs) + zI0, zR0, zI, zR = self.forward(y_cond, orig=True) # (T, nodes, samples) + zI = zI.clone().detach().requires_grad_(True); zI.retain_grad() + zR = zR.clone().detach().requires_grad_(True); zR.retain_grad() + lik = torch.zeros(n_samples, dtype=zI.dtype, device=self.device) + lI1, lI0 = F.logsigmoid(zI), F.logsigmoid(-zI) + lR1, lR0 = F.logsigmoid(zR), F.logsigmoid(-zR) + for i in range(K): + TL, TR = obs_time[i - 1] if i > 0 else None, obs_time[i] + if TL is None: + for t in range(TR - 1, -1, -1): + lik = lik + self._lik_step(Y[t], Y[t + 1], lI1[t], lI0[t], lR1[t], lR0[t]) + else: + yL = Y[TL] # (nodes, samples) + L = TR - TL + reach = torch.zeros(L, self.n_nodes, n_samples, dtype=torch.bool, device=self.device) + reach[0] = (yL == SIR_STATES.I) # (nodes, samples) + yL_not_R = (yL != SIR_STATES.R) # (nodes, samples) + for d in range(1, L): + reach[d] = reach[d - 1] | ((torch.sparse.mm(self.adj, reach[d - 1].float()) > 0) & yL_not_R) + for t in range(TR - 1, TL, -1): + lik = lik + self._lik_step(Y[t], Y[t + 1], lI1[t], lI0[t], lR1[t], lR0[t], yL = yL, reach = reach[t - TL]) + return lik, zI0, zR0, zI, zR + @torch.no_grad() + def clamp_grad(self, z0, grad): + return torch.where(z0 < self.zlim, torch.where(z0 > -self.zlim, grad, F.relu(grad)), -F.relu(-grad)) + def backward(self, loss, zI0, zR0, zI, zR): + loss.backward() + z0 = torch.stack([zI0, zR0], dim = 0) + z0.backward(torch.stack([self.clamp_grad(zI0, zI.grad), self.clamp_grad(zR0, zR.grad)], dim = 0)) + @torch.no_grad() + def _samp_step(self, y, zI, zR, compute_lik=False, yL=None, reach=None): + if reach is not None:# y: (nodes, samples); yL: (nodes,); z: (nodes, 1); reach: (nodes,) + reach = reach.unsqueeze(dim=1) # (nodes, 1) + n_samples = y.size(dim=1) + if compute_lik: + lik = self.zero + # unreachable + if reach is not None: + y = torch.where(reach, y, yL.unsqueeze(dim=1)) + # R->I + qR = torch.sigmoid(zR) # (nodes, 1) + xR = SIR_STATES.R - qR.expand(-1, n_samples).bernoulli().long() # (nodes, samples) + msk = (y == SIR_STATES.R) # (nodes, samples) + if yL is not None: + msk = msk & (yL != SIR_STATES.R).unsqueeze(dim=1) & reach + y = torch.where(msk, xR, y) + if compute_lik: + lik = lik + torch.where(msk, torch_log(torch.where(xR != SIR_STATES.R, qR, 1. - qR)), self.zero).sum(dim=0) # (samples,) + # I->S + zI, uid = zI.sort(dim = 0, descending = True) + uid = uid.squeeze(dim = -1) # (nodes,) + qI = torch.sigmoid(zI) # (nodes, 1) (already sorted by zI) + xI = SIR_STATES.I - qI.expand(-1, n_samples).bernoulli().long() # (nodes, samples) + if compute_lik: + lI = torch_log(torch.where(xI != SIR_STATES.I, qI, 1. - qI)) # (nodes, samples) + msk = (y == SIR_STATES.I) # (nodes, samples) + rem = torch.where(msk, (reach.long() if reach is not None else 1) + torch.sparse.mm(self.adj, msk.float()).long(), self.n_inf) # (nodes, samples) + for i, u in enumerate(uid): + if msk[u].max(): + vid = self.neighbs[u.item()] # (neighbs,) + rem_u = rem[u] # (samples,) + rem_v = rem[vid] # (neighbs, samples) + opt = (rem_u > 1) & (rem_v.min(dim=0).values > 1) # (samples,) + if yL is not None: + opt = opt & (yL[u] != SIR_STATES.I) + msk_opt = msk[u] & opt + y[u] = torch.where(msk_opt, xI[i], y[u]) # (samples,) + trs = (y[u] != SIR_STATES.I) # (samples,) + rem[vid] = torch.where(msk[u].unsqueeze(dim=0), torch.where(trs.unsqueeze(dim=0), rem_v - 1, rem_v), rem_v) + rem[u] = torch.where(msk[u], torch.where(trs, rem_u - 1, self.n_inf), rem[u]) + msk[u] = msk_opt + if compute_lik: + lik = lik + torch.where(msk[uid], lI, self.zero).sum(dim=0) # (samples,) + return y, lik + else: + return y, 0. + @torch.no_grad() + def _samp_seg(self, yR, zI, zR, n_samples, TL, TR, yL=None, compute_lik=False): # yR: (nodes,); zI: (T, nodes, 1); zR: (T, nodes, 1); return: Y: (TR - TL, nodes, samples), lik: (samples,) + if compute_lik: # segment log-likelihood under the proposal Q_theta + lik = torch.zeros(n_samples, dtype=torch.float, device=self.device) + y = yR.unsqueeze(dim=1).expand(-1, n_samples) # (nodes, samples) + if yL is None: # no left constraint, sample purely by original backward local-support steps + Y = torch.empty(TR, self.n_nodes, n_samples, dtype=torch.long, device=self.device) + for t in range(TR - 1, -1, -1): + y, lik_t = self._samp_step(y, zI[t], zR[t], compute_lik=compute_lik) + if compute_lik: + lik = lik + lik_t + Y[t] = y + else: + L = TR - TL + Y = torch.empty(L, self.n_nodes, n_samples, dtype=torch.long, device=self.device) + Y[0] = yL.unsqueeze(dim = -1) + reach = torch.zeros(L, self.n_nodes, dtype=torch.bool, device=yL.device) + reach[0] = (yL == SIR_STATES.I) + + yL_not_R = (yL != SIR_STATES.R) # (nodes,) FIX: keep as 1D boolean + for d in range(1, L): + nbr = (torch.sparse.mm(self.adj, reach[d - 1].float().unsqueeze(1)).squeeze(1) > 0) # (nodes,) + reach[d] = reach[d - 1] | (nbr & yL_not_R) + for t in range(TR - 1, TL, -1): + y, lik_t = self._samp_step(y, zI[t], zR[t], compute_lik=compute_lik, yL = yL, reach = reach[t - TL]) + if compute_lik: + lik = lik + lik_t + Y[t - TL] = y + if compute_lik: + return Y.detach().clone(), lik.detach().clone() + else: + return Y.detach().clone(), 0. + @torch.no_grad() + def samp_ms(self, y, zI, zR, n_samples, obs_time, compute_lik=False): # z: (T, nodes, 1) + # y: (nodes, T+1) + Y = torch.empty(self.T, self.n_nodes, n_samples, dtype=torch.long, device=self.device) + for t in obs_time: + if t < self.T: + Y[t] = y[:, t].unsqueeze(dim=1).expand(-1, n_samples) + if compute_lik: + lik = 0. # will become (samples,) after first addition + for i in range(len(obs_time) - 1, -1, -1): # Sample segments in reverse order (right endpoint always known). + TL, TR = obs_time[i - 1] if i > 0 else None, obs_time[i] + yL = None if TL is None else y[:, TL] + yR = y[:, TR] + Y[TL : TR], lik_seg = self._samp_seg(yR, zI, zR, n_samples, TL, TR, yL = yL, compute_lik = compute_lik) + if compute_lik: + lik = lik + lik_seg + if compute_lik: + return Y.detach().clone(), lik.detach().clone() + else: + return Y.detach().clone() + + +def q_loss(q_net, data, I0, bpar, n_samples, obs_time): + T = data.T.item() + n_nodes = data.num_nodes + Y = diffus_gen( + T=T, + n_nodes=n_nodes, + edge_index=data.edge_index, + I0=I0, + n_samples=n_samples, + pI=bpar.pI, + pR=bpar.pR, + ) # (T+1, nodes, samples) + + q_liks, zI0, zR0, zI, zR = q_net.lik_ms(Y=Y, obs_time=obs_time) + return -q_liks.mean(), zI0, zR0, zI, zR + +def q_train(data, obs_time, bpar, args): + I0 = (data.y[:, 0] == 1).long().sum().item() + q_net = QNet.make(data, obs_time, args) + q_net.train() + opt = optim.AdamW(q_net.parameters(), lr=args.q_lr) + pbar = trange(1, args.q_steps + 1) + for step in pbar: + opt.zero_grad() + loss, zI0, zR0, zI, zR = q_loss(q_net, data, I0, bpar, args.q_samples, obs_time=obs_time) + pbar.set_description(f'[step={step}] loss={loss.item():.4f}') + q_net.backward(loss, zI0, zR0, zI, zR) + opt.step() + q_net.eval() + return q_net + +@torch.no_grad() +def t_mcmc(data, bpar, q_net, args, obs_time, keepdim=True): + + I0 = (data.y[:, 0] == 1).long().sum().item() + + obs_time = sorted(list(obs_time)) + y_obs = torch.stack([data.y[:, t : t + 1] for t in obs_time], dim=2) # (nodes, 1, obs) + zI, zR = q_net(y_obs) # (T, nodes, 1) + + X, lqX = q_net.samp_ms(data.y, zI, zR, args.t_samples, obs_time=obs_time, compute_lik=True) + lpX = diffus_liks(Y=X, edge_index=data.edge_index, I0=I0, coef=args.p_coef, pI=bpar.pI, pR=bpar.pR) + + tI_avg = data_make_t(X, SIR_STATES.I, dim=0).float().mean(dim=1, keepdim=keepdim) + tR_avg = data_make_t(X, SIR_STATES.R, dim=0).float().mean(dim=1, keepdim=keepdim) + + pbar = trange(1, args.t_steps + 1) + for step in pbar: + Y, lqY = q_net.samp_ms(data.y, zI, zR, args.t_samples, obs_time=obs_time, compute_lik=True) + lpY = diffus_liks(Y=Y, edge_index=data.edge_index, I0=I0, coef=args.p_coef, pI=bpar.pI, pR=bpar.pR) + + # Hastings acceptance + a = torch.rand(args.t_samples, device=args.device) <= torch.exp(lpY + lqX - lpX - lqY) + + X = torch.where(a, Y, X) + lqX = torch.where(a, lqY, lqX) + lpX = torch.where(a, lpY, lpX) + + tI = data_make_t(X, SIR_STATES.I, dim=0).float().mean(dim=1, keepdim=keepdim) + tR = data_make_t(X, SIR_STATES.R, dim=0).float().mean(dim=1, keepdim=keepdim) + tI_avg = args.t_keep * tI_avg + (1.0 - args.t_keep) * tI + tR_avg = args.t_keep * tR_avg + (1.0 - args.t_keep) * tR + + return tI_avg, tR_avg + +def main(data): + # parse obs times + obs_time = [int(t) for t in args.obs_time.split(',') if t] + obs_time.append(data.T.item()) + obs_time = sorted(set(obs_time)) + # estimate diffusion parameters + bpar = b_estim(data, args) + print(f'[est] pI={bpar.pI:.4f}, pR={bpar.pR:.4f}', flush = True) + # train a proposal network + q_net = q_train(data, obs_time, bpar, args) + # estimate transition times + tI, tR = t_mcmc(data, bpar, q_net, args, obs_time = obs_time, keepdim = True) # (nodes, 1) + T = data.T.item() + tI = tI.round().long() + tR = tR.round().long() + # compose a history + with torch.no_grad(): + y_pred = torch.zeros_like(data.y) # (nodes, T+1) + y_pred.scatter_(dim = 1, index = torch.minimum(tI, data.T), src = torch.full_like(tI, 1)) + y_pred.scatter_(dim = 1, index = torch.minimum(tR, data.T), src = torch.full_like(tR, 2)) + y_pred = y_pred[:, : data.T.item()].cummax(dim = 1).values + return y_pred + +args = get_args() +tester = Tester(args.data_dir, args.device, main) +tester.test([args.dataset], seed = args.seed, rep = 1) +tester.save(args.output) \ No newline at end of file diff --git a/inc/data.py b/inc/data.py index 2e10d29..41a2a8e 100644 --- a/inc/data.py +++ b/inc/data.py @@ -66,7 +66,7 @@ def data_synthetic(graph, diffus, data_dir, device): data_dir = osp.join(data_dir, 'synthetic') f_data = osp.join(data_dir, f'{graph}-{diffus}.pt') if osp.exists(f_data): - return torch.load(f_data, map_location = device) + return torch.load(f_data, map_location=device, weights_only=False) else: seed = 123456789 T = 10 @@ -88,7 +88,7 @@ def data_prost(diffus, data_dir, device): data_dir = osp.join(data_dir, 'prost') f_data = osp.join(data_dir, f'prost-{diffus}.pt') if osp.exists(f_data): - return torch.load(f_data, map_location = device) + return torch.load(f_data, map_location=device, weights_only=False) else: seed = 123456789 T = 15 @@ -110,7 +110,7 @@ def data_oregon2(diffus, data_dir, device): data_dir = osp.join(data_dir, 'oregon2') f_data = osp.join(data_dir, f'oregon2-{diffus}.pt') if osp.exists(f_data): - return torch.load(f_data, map_location = device) + return torch.load(f_data, map_location=device, weights_only=False) else: seed = 123456789 T = 15 @@ -130,7 +130,7 @@ def data_farmers_si(data_dir, device): data_dir = osp.join(data_dir, 'farmers') f_data = osp.join(data_dir, 'farmers-si.pt') if osp.exists(f_data): - return torch.load(f_data, map_location = device) + return torch.load(f_data, map_location=device, weights_only=False) else: f_raw = file_require(None, data_dir, 'brfarmers.rdata') df = pyreadr.read_r('farmers/brfarmers.rdata')['brfarmers'] @@ -160,7 +160,7 @@ def data_pol_si(data_dir, device): data_dir = osp.join(data_dir, 'pol') f_data = osp.join(data_dir, 'pol-si.pt') if osp.exists(f_data): - return torch.load(f_data, map_location = device) + return torch.load(f_data, map_location=device, weights_only=False) else: f_edge = file_require('https://nrvis.com/download/data/rt/rt-pol.zip', data_dir, 'rt-pol.txt', z = 'zip') df_fr, df_to, df_time = [], [], [] @@ -189,7 +189,7 @@ def data_covid_sir(data_dir, device): data_dir = osp.join(data_dir, 'covid') f_data = osp.join(data_dir, 'covid-sir.pt') if osp.exists(f_data): - return torch.load(f_data, map_location=device) + return torch.load(f_data, map_location=device, weights_only=False) else: COVID_KNN = 10 f_s2a = file_require(None, data_dir, 'state2abbr.pyon') @@ -250,7 +250,7 @@ def data_heb_sir(data_dir, device): data_dir = osp.join(data_dir, 'heb') f_data = osp.join(data_dir, 'heb-sir.pt') if osp.exists(f_data): - return torch.load(f_data, map_location = device) + return torch.load(f_data, map_location=device, weights_only=False) else: f_edge = file_require(url = None, fdir = data_dir, fname = 'DS1_NON_VIRAL_Gtw.tsv') df = pd.read_csv(f_edge, sep = '\t', header = None, names = ['time', 'to', 'fr'], dtype = dict(time = str, to = int, fr = int), parse_dates = ['time']) diff --git a/inc/diffus.py b/inc/diffus.py index 40dc77f..cc7d7f6 100644 --- a/inc/diffus.py +++ b/inc/diffus.py @@ -62,50 +62,94 @@ def __repr__(self, digits = 4): def dict(self): return Dict(pI = self.pI.item(), pR = self.pR.item()) -def b_lik(bpar, data): # yT: (nodes,) +def _mf_init_from_prior(data, n_nodes, device): + # prior: only use I0 count (same as current code) + I0 = (data.y[:, 0] == SIR_STATES.I).sum() + lI = torch.full((n_nodes,), I0 / n_nodes, dtype=torch.float, device=device) + lS = torch.full((n_nodes,), 1. - I0 / n_nodes, dtype=torch.float, device=device) + lR = torch.zeros(n_nodes, dtype=torch.float, device=device) + return lS, lI, lR + +def _mf_init_from_snapshot(y, device): + # hard clamp to observed snapshot y (nodes,) + lS = (y == SIR_STATES.S).float().to(device) + lI = (y == SIR_STATES.I).float().to(device) + lR = (y == SIR_STATES.R).float().to(device) + return lS, lI, lR + +def b_lik(bpar, data, obs_time=None): device = data.y.device n_nodes = data.num_nodes T = data.T.item() ei = data.edge_index - I0 = (data.y[:, 0] == 1).sum() pI, pR = bpar.pI, bpar.pR - lSs, lIs, lRs = [], [], [] - lIs.append(torch.full((n_nodes,), I0 / n_nodes, dtype = torch.float, device = device)) - lSs.append(torch.full((n_nodes,), 1. - I0 / n_nodes, dtype = torch.float, device = device)) - lRs.append(torch.zeros(n_nodes, dtype = torch.float, device = device)) - for t in range(T): - aI = pysc.scatter_mul(src = (1. - lIs[-1] * pI)[ei[0]], dim = 0, index = ei[1], dim_size = n_nodes) - lS = lSs[-1] * aI - kI = lIs[-1] + lSs[-1] * (1. - aI) - lI = kI * (1. - pR) - lR = lRs[-1] + kI * pR - lSs.append(lS) - lIs.append(lI) - lRs.append(lR) - lik = torch.stack([lSs[-1], lIs[-1], lRs[-1]], dim = 0) # (states, nodes) - yT = data.y[:, -1].unsqueeze(dim = 0) # (1, nodes) - lik = lik.gather(dim = 0, index = yT) # (1, nodes) - lik = torch_log(lik).mean() - return lik - -def b_estim(data, args): + + # -------- obs times -------- + if obs_time is None: + obs_time = [T] + obs_time = sorted(set(int(t) for t in obs_time if 0 <= int(t) <= T)) + if len(obs_time) == 0 or obs_time[-1] != T: + obs_time.append(T) + + # -------- segmented mean-field -------- + lS, lI, lR = _mf_init_from_prior(data, n_nodes, device) + t_prev = 0 + lik_total = 0.0 + + for t_obs in obs_time: + # forward from t_prev -> t_obs + for _ in range(t_obs - t_prev): + aI = pysc.scatter_mul( + src=(1. - lI * pI)[ei[0]], + dim=0, + index=ei[1], + dim_size=n_nodes + ) + lS_new = lS * aI + kI = lI + lS * (1. - aI) + lI_new = kI * (1. - pR) + lR_new = lR + kI * pR + lS, lI, lR = lS_new, lI_new, lR_new + + # score snapshot at t_obs + lik = torch.stack([lS, lI, lR], dim=0) # (3, nodes) + y = data.y[:, t_obs].unsqueeze(dim=0) # (1, nodes) + prob = lik.gather(dim=0, index=y).squeeze(dim=0) # (nodes,) + lik_total = lik_total + torch_log(prob).mean() + + # clamp for next segment (if any) + lS, lI, lR = _mf_init_from_snapshot(data.y[:, t_obs], device) + t_prev = t_obs + + # optional: normalize by number of observed frames, keeps loss scale stable + lik_total = lik_total / len(obs_time) + return lik_total + +def b_estim(data, args, obs_time=None): T = data.T.item() - n_nodes = data.num_nodes - n_edges = data.edge_index.size(dim = 1) n_cls = data.y[:, T].max().item() + 1 pI, pR = args.b_pI0, (args.b_pR0 if n_cls == 3 else 0.) - #print(f'[ini] pI={pI:.4f}, pR={pR:.4f}', flush = True) device = data.y.device - bpar = BPar(pI = pI, pR = pR, device = device) + + # if caller doesn't pass obs_time, parse from args (compatible with current main) + if obs_time is None and hasattr(args, "obs_time"): + tmp = [int(t) for t in str(args.obs_time).split(',') if t] + tmp.append(T) + obs_time = tmp + + bpar = BPar(pI=pI, pR=pR, device=device) bpar.train() - opt = optim.AdamW(bpar.parameters(), lr = args.b_lr, betas = (0.5, 0.5)) + opt = optim.AdamW(bpar.parameters(), lr=args.b_lr, betas=(0.5, 0.5)) pbar = trange(1, args.b_steps + 1) + for step in pbar: opt.zero_grad() - loss = -b_lik(bpar, data) + loss = -b_lik(bpar, data, obs_time=obs_time) loss.backward() opt.step() bpar.clamp_() pbar.set_description(f'[step={step}] {bpar}') + bpar.eval() return bpar.dict() + diff --git a/inc/test_ms.py b/inc/test_ms.py index 8aa60af..61a1855 100644 --- a/inc/test_ms.py +++ b/inc/test_ms.py @@ -58,13 +58,13 @@ def test_skm(skm_fn, data, y_pred, **kwargs): @torch.no_grad() def test_nrmse(data, y_pred): - y_fixed = test_fix_obs(data, y_pred) - tI_pred = data_make_tI(y_fixed, dim=-1) + y_pred = test_fix_obs(data, y_pred) + tI_pred = data_make_t(y_pred, SIR_STATES.I, dim = -1) mse = skm.mean_squared_error(torch2np(data.tI), torch2np(tI_pred)) if hasattr(data, 'tR'): - tR_pred = data_make_t(y_fixed, SIR_STATES.R, dim=-1) + tR_pred = data_make_t(y_pred, SIR_STATES.R, dim = -1) mseR = skm.mean_squared_error(torch2np(data.tR), torch2np(tR_pred)) - mse = (mse + mseR) / 2.0 + mse = (mse + mseR) / 2. nrmse = np.sqrt(mse) / (data.T.item() + 1) return float(nrmse) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..8a56907 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,65 @@ +aiohappyeyeballs==2.6.1 +aiohttp==3.13.3 +aiosignal==1.4.0 +attrs==25.4.0 +bokeh==3.8.2 +certifi==2026.1.4 +charset-normalizer==3.4.4 +colorama==0.4.6 +contourpy==1.3.3 +cycler==0.12.1 +decorator==5.2.1 +dynetx==0.3.2 +filelock==3.20.0 +fonttools==4.61.1 +frozenlist==1.8.0 +fsspec==2025.12.0 +future==1.0.0 +idna==3.11 +igraph==1.0.0 +Jinja2==3.1.6 +joblib==1.5.3 +kiwisolver==1.4.9 +MarkupSafe==2.1.5 +matplotlib==3.10.8 +mpmath==1.3.0 +multidict==6.7.0 +narwhals==2.15.0 +ndlib==5.1.1 +netdispatch==0.1.0 +networkx==3.6.1 +numpy==2.3.5 +packaging==26.0 +pandas==3.0.0 +pillow==12.0.0 +propcache==0.4.1 +psutil==7.2.1 +pyg-lib==0.5.0+pt28cu128 +pyparsing==3.3.2 +python-dateutil==2.9.0.post0 +python-igraph==1.0.0 +PyYAML==6.0.3 +requests==2.32.5 +scikit-learn==1.8.0 +scipy==1.17.0 +seaborn==0.13.2 +six==1.17.0 +sympy==1.14.0 +texttable==1.7.0 +threadpoolctl==3.6.0 +torch==2.8.0+cu128 +torch-geometric==2.7.0 +torch_cluster==1.6.3+pt28cu128 +torch_scatter==2.1.2+pt28cu128 +torch_sparse==0.6.18+pt28cu128 +torch_spline_conv==1.2.2+pt28cu128 +torchaudio==2.8.0+cu128 +torchvision==0.23.0+cu128 +tornado==6.5.4 +tqdm==4.67.1 +typing_extensions==4.15.0 +tzdata==2025.3 +urllib3==2.6.3 +xxhash==3.6.0 +xyzservices==2025.11.0 +yarl==1.22.0 From 02fee1a9fb5421a99d2f6ced61d57b42fc2e3575 Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Tue, 27 Jan 2026 13:36:22 -0600 Subject: [PATCH 05/19] update hermes for OOM issues --- hermes.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/hermes.py b/hermes.py index bb99f34..44b6921 100644 --- a/hermes.py +++ b/hermes.py @@ -115,10 +115,11 @@ def _lik_step(self, y0, y1, lI1, lI0, lR1, lR0, yL=None, reach=None): # y*, l*, msk_opt = msk_u & opt # (samples,) lik = lik + torch.where(msk_opt, torch.where(y0 == SIR_STATES.S, lI1, lI0), self.zero) trs = (y0.gather(dim = 0, index = u.unsqueeze(dim=0)).squeeze(dim=0) == SIR_STATES.S) # (samples,) - rem = rem - torch.zeros_like(rem).index_put( - (self.eidx[1, :, None].expand(-1, n_samples).flatten(), cols[None].expand(self.n_edges, -1).flatten()), - ((msk_u & trs) & (self.eidx[0, :, None] == u)).flatten().long(), # (edges * samples,) - accumulate = False) + rem = rem - (torch.sparse.mm(self.adj, torch.zeros(rem.size(), dtype=self.adj.dtype, device=rem.device).scatter(dim = 0, index = u.unsqueeze(dim=0), src = (msk_u & trs).unsqueeze(dim=0).to(self.adj.dtype))) > 0).to(rem.dtype) # (nodes, samples) + # rem = rem - torch.zeros_like(rem).index_put( + # (self.eidx[1, :, None].expand(-1, n_samples).flatten(), cols[None].expand(self.n_edges, -1).flatten()), + # ((msk_u & trs) & (self.eidx[0, :, None] == u)).flatten().long(), # (edges * samples,) + # accumulate = False) rem = rem.index_put((u, cols), torch.where(msk_u, torch.where(trs, rem_u - 1, self.n_inf), rem.gather(dim = 0, index = u.unsqueeze(dim=0)).squeeze(dim=0))) msk = msk.index_put((u, cols), msk_opt) return lik From 48020d6001999fa8cfae88b9e200f178ccc5585e Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Wed, 28 Jan 2026 21:59:44 -0600 Subject: [PATCH 06/19] update 3 new baselines --- brits.py | 401 +++++++++++++++++++++++++++++ brits_ms.py | 537 +++++++++++++++++++++++++++++++++++++++ grin.py | 219 ++++++++++++++++ grin_ms.py | 272 ++++++++++++++++++++ spin.py | 672 +++++++++++++++++++++++++++++++++++++++++++++++++ spin_ms.py | 709 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 6 files changed, 2810 insertions(+) create mode 100644 brits.py create mode 100644 brits_ms.py create mode 100644 grin.py create mode 100644 grin_ms.py create mode 100644 spin.py create mode 100644 spin_ms.py diff --git a/brits.py b/brits.py new file mode 100644 index 0000000..a91724b --- /dev/null +++ b/brits.py @@ -0,0 +1,401 @@ +# ! pip install class-resolver==0.3.10 +# ! pip install --no-index torch-scatter==2.0.7 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +# ! pip install --no-index torch-sparse==0.6.9 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +# ! pip install --no-index torch-cluster==1.5.9 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +# ! pip install --no-index torch-spline-conv==1.2.1 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +# ! pip install torch-geometric==2.0.4 +# ! pip install ndlib==5.1.1 + +from inc.diffus import * +from inc.test import * + +import argparse +import torch + + +def get_args(argv=None): + """Parse command-line arguments. + + We keep the original notebook defaults so running without extra flags + behaves the same as before. + """ + parser = argparse.ArgumentParser() + + # ---- experiment I/O ---- + parser.add_argument('--dataset', type=str, required=True, help='dataset name') + parser.add_argument('--seed', type=int, default=123456789, help='random seed') + parser.add_argument('--data_dir', type=str, default='input', help='dataset folder') + parser.add_argument('--output', type=str, default='output/brits.pt', help='output file name') + parser.add_argument( + '--device', + type=torch.device, + default=torch.device('cuda' if torch.cuda.is_available() else 'cpu'), + help='torch device, e.g., cpu, cuda, cuda:0' + ) + + # ---- diffusion parameter estimation (b_*) ---- + parser.add_argument('--b_pI0', type=float, default=1e-3, + help='initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type=float, default=1e-3, + help='initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type=int, default=500, + help='optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type=float, default=3e-3, + help='learning rate in diffusion parameter estimation') + + # ---- BRITS hyperparameters ---- + parser.add_argument('--lr', type=float, default=1e-3, help='learning rate') + parser.add_argument('--epochs', type=int, default=1000, help='training epochs') + parser.add_argument('--batch_size', type=int, default=64, help='batch size') + parser.add_argument('--hid_size', type=int, default=108, help='RNN hidden size') + parser.add_argument('--impute_weight', type=float, default=0.3, help='imputation loss weight') + parser.add_argument('--label_weight', type=float, default=1.0, help='label loss weight') + + if argv is None: + return parser.parse_args() + return parser.parse_args(argv) + + +'''https://github.com/caow13/BRITS''' +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.optim as optim + +from torch.autograd import Variable +from torch.nn.parameter import Parameter + +import math +# import utils +import argparse +# import data_loader + +# from ipdb import set_trace +from sklearn import metrics + + +def binary_cross_entropy_with_logits(input, target, weight=None, size_average=True, reduce=True): + if not (target.size() == input.size()): + raise ValueError("Target size ({}) must be the same as input size ({})".format(target.size(), input.size())) + max_val = (-input).clamp(min=0) + loss = input - input * target + max_val + ((-max_val).exp() + (-input - max_val).exp()).log() + if weight is not None: + loss = loss * weight + if not reduce: + return loss + elif size_average: + return loss.mean() + else: + return loss.sum() + + +class FeatureRegression(nn.Module): + def __init__(self, input_size): + super().__init__() + self.build(input_size) + + def build(self, input_size): + self.W = Parameter(torch.Tensor(input_size, input_size)) + self.b = Parameter(torch.Tensor(input_size)) + m = torch.ones(input_size, input_size) - torch.eye(input_size, input_size) + self.register_buffer('m', m) + self.reset_parameters() + + def reset_parameters(self): + stdv = 1. / math.sqrt(self.W.size(0)) + self.W.data.uniform_(-stdv, stdv) + if self.b is not None: + self.b.data.uniform_(-stdv, stdv) + + def forward(self, x): + z_h = F.linear(x, self.W * Variable(self.m), self.b) + return z_h + + +class TemporalDecay(nn.Module): + def __init__(self, input_size, output_size, diag=False): + super().__init__() + self.diag = diag + self.build(input_size, output_size) + + def build(self, input_size, output_size): + self.W = Parameter(torch.Tensor(output_size, input_size)) + self.b = Parameter(torch.Tensor(output_size)) + if self.diag == True: + assert (input_size == output_size) + m = torch.eye(input_size, input_size) + self.register_buffer('m', m) + self.reset_parameters() + + def reset_parameters(self): + stdv = 1. / math.sqrt(self.W.size(0)) + self.W.data.uniform_(-stdv, stdv) + if self.b is not None: + self.b.data.uniform_(-stdv, stdv) + + def forward(self, d): + if self.diag == True: + gamma = F.relu(F.linear(d, self.W * Variable(self.m), self.b)) + else: + gamma = F.relu(F.linear(d, self.W, self.b)) + gamma = torch.exp(-gamma) + return gamma + + +class RITS(nn.Module): + def __init__(self, xdim, rnn_hid_size, impute_weight, label_weight): + super().__init__() + self.xdim = xdim + self.rnn_hid_size = rnn_hid_size + self.impute_weight = impute_weight + self.label_weight = label_weight + self.build() + + def build(self): + self.rnn_cell = nn.LSTMCell(self.xdim * 2, self.rnn_hid_size) + self.temp_decay_h = TemporalDecay(input_size=self.xdim, output_size=self.rnn_hid_size, diag=False) + self.temp_decay_x = TemporalDecay(input_size=self.xdim, output_size=self.xdim, diag=True) + self.hist_reg = nn.Linear(self.rnn_hid_size, self.xdim) + self.feat_reg = FeatureRegression(self.xdim) + self.weight_combine = nn.Linear(self.xdim * 2, self.xdim) + self.dropout = nn.Dropout(p=0.25) + self.out = nn.Linear(self.rnn_hid_size, 1) + + def forward(self, data, direct): + values = data[direct]['values'] + masks = data[direct]['masks'] + deltas = data[direct]['deltas'] + evals = data[direct]['evals'] + eval_masks = data[direct]['eval_masks'] + labels = data['labels'].reshape((-1, 1)) + is_train = data['is_train'].reshape((-1, 1)) + h = Variable(torch.zeros((values.size(0), self.rnn_hid_size))) + c = Variable(torch.zeros((values.size(0), self.rnn_hid_size))) + if torch.cuda.is_available(): + h, c = h.cuda(), c.cuda() + x_loss = 0.0 + y_loss = 0.0 + imputations = [] + for t in range(min(values.size(1), masks.size(1), deltas.size(1))): + x = values[:, t, :] + m = masks[:, t, :] + d = deltas[:, t, :] + gamma_h = self.temp_decay_h(d) + gamma_x = self.temp_decay_x(d) + h = h * gamma_h + x_h = self.hist_reg(h) + x_loss += torch.sum(torch.abs(x - x_h) * m) / (torch.sum(m) + 1e-5) + x_c = m * x + (1 - m) * x_h + z_h = self.feat_reg(x_c) + x_loss += torch.sum(torch.abs(x - z_h) * m) / (torch.sum(m) + 1e-5) + alpha = self.weight_combine(torch.cat([gamma_x, m], dim=1)) + c_h = alpha * z_h + (1 - alpha) * x_h + x_loss += torch.sum(torch.abs(x - c_h) * m) / (torch.sum(m) + 1e-5) + c_c = m * x + (1 - m) * c_h + inputs = torch.cat([c_c, m], dim=1) + h, c = self.rnn_cell(inputs, (h, c)) + imputations.append(c_c.unsqueeze(dim=1)) + imputations = torch.cat(imputations, dim=1) + y_h = self.out(h) + y_loss = binary_cross_entropy_with_logits(y_h, labels, reduce=False) + y_loss = torch.sum(y_loss * is_train) / (torch.sum(is_train) + 1e-5) + y_h = torch.sigmoid(y_h) + return {'loss': x_loss * self.impute_weight + y_loss * self.label_weight, 'predictions': y_h, \ + 'imputations': imputations, 'labels': labels, 'is_train': is_train, \ + 'evals': evals, 'eval_masks': eval_masks} + + def run_on_batch(self, data, optimizer, epoch=None): + ret = self(data, direct='forward') + if optimizer is not None: + optimizer.zero_grad() + ret['loss'].backward() + optimizer.step() + return ret + + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.optim as optim + +from torch.autograd import Variable +from torch.nn.parameter import Parameter + +import math +# import utils +import argparse +# import data_loader + +# import rits +from sklearn import metrics + + +# from ipdb import set_trace + +class BRITS(nn.Module): + def __init__(self, xdim, rnn_hid_size, impute_weight, label_weight): + super().__init__() + self.xdim = xdim + self.rnn_hid_size = rnn_hid_size + self.impute_weight = impute_weight + self.label_weight = label_weight + self.build() + + def build(self): + self.rits_f = RITS(self.xdim, self.rnn_hid_size, self.impute_weight, self.label_weight) + self.rits_b = RITS(self.xdim, self.rnn_hid_size, self.impute_weight, self.label_weight) + + def forward(self, data): + ret_f = self.rits_f(data, 'forward') + ret_b = self.reverse(self.rits_b(data, 'backward')) + ret = self.merge_ret(ret_f, ret_b) + return ret + + def merge_ret(self, ret_f, ret_b): + loss_f = ret_f['loss'] + loss_b = ret_b['loss'] + loss_c = self.get_consistency_loss(ret_f['imputations'], ret_b['imputations']) + loss = loss_f + loss_b + loss_c + predictions = (ret_f['predictions'] + ret_b['predictions']) / 2 + imputations = (ret_f['imputations'] + ret_b['imputations']) / 2 + ret_f['loss'] = loss + ret_f['predictions'] = predictions + ret_f['imputations'] = imputations + return ret_f + + def get_consistency_loss(self, pred_f, pred_b): + loss = torch.abs(pred_f - pred_b).mean() * 1e-1 + return loss + + def reverse(self, ret): + def reverse_tensor(tensor_): + if tensor_.dim() <= 1: + return tensor_ + indices = range(tensor_.size()[1])[::-1] + indices = Variable(torch.LongTensor(indices), requires_grad=False) + if torch.cuda.is_available(): + indices = indices.cuda() + return tensor_.index_select(1, indices) + + for key in ret: + ret[key] = reverse_tensor(ret[key]) + return ret + + def run_on_batch(self, data, optimizer, epoch=None): + ret = self(data) + if optimizer is not None: + optimizer.zero_grad() + ret['loss'].backward() + optimizer.step() + return ret + + +import copy +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.optim as optim +from torch.optim.lr_scheduler import StepLR + +import numpy as np + +import time +# import utils +# import models +import argparse +# import data_loader +import pandas as pd +import ujson as json + +from sklearn import metrics + + +# from ipdb import set_trace + +def to_var(var): + if torch.is_tensor(var): + device = var.device + var = torch.autograd.Variable(var) + var = var.to(device) + return var + if isinstance(var, int) or isinstance(var, float) or isinstance(var, str): + return var + if isinstance(var, dict): + return {key: to_var(val) for key, val in var.items()} + if isinstance(var, list): + return [to_var(x) for x in var] + + +@torch.no_grad() +def brits_prep_rec(y, back): # (samples, T + 1, nodes) + evals = y.float().clone() + values = evals.clone() + if back: + values[:, 1:] = 0 + else: + values[:, : -1] = 0 + masks = torch.zeros_like(y) + if back: + masks[:, 0] = True + else: + masks[:, -1] = True + eval_masks = masks.clone() + deltas = torch.cat([torch.arange(y.size(dim=1) - 1, dtype=torch.float, device=y.device), + torch.zeros(1, dtype=torch.float, device=y.device)], dim=0).unsqueeze(dim=-1).expand(*y.size()) + return dict(values=values.contiguous(), masks=masks.contiguous(), evals=evals.contiguous(), + eval_masks=eval_masks.contiguous(), deltas=deltas.contiguous()) + + +@torch.no_grad() +def brits_prep(y, is_train): # y: (samples, nodes, T + 1) + n_samples = y.size(dim=0) + y = y.transpose(1, 2) # (samples, T + 1, nodes) + return to_var(dict( + forward=brits_prep_rec(y, back=False), + backward=brits_prep_rec(y.flip(dims=[1]), back=True), + labels=torch.zeros(n_samples, 1, dtype=torch.long, device=y.device), + is_train=torch.tensor([is_train] * n_samples, dtype=torch.float, device=y.device), + )) + + +def brits_run(data, args): + """Run BRITS on a single dataset instance (loaded by Tester).""" + bpar = b_estim(data, args) + + model = BRITS(data.num_nodes, args.hid_size, args.impute_weight, args.label_weight) + model = model.to(args.device) + + # train + I0 = (data.y[:, 0] == SIR_STATES.I).long().sum().item() + optimizer = optim.Adam(model.parameters(), lr=args.lr) + pbar = trange(args.epochs) + for epoch in pbar: + model.train() + batch = brits_prep(diffus_gen(T=data.T.item(), n_nodes=data.num_nodes, edge_index=data.edge_index, I0=I0, + n_samples=args.batch_size, pI=bpar.pI, pR=bpar.pR).transpose(0, 2).clone(), + is_train=1) + ret = model.run_on_batch(batch, optimizer, epoch) + pbar.set_description(f'epoch={epoch + 1} loss={ret["loss"].item():.4f}') + + # infer + with torch.no_grad(): + model.eval() + rec = brits_prep(data.y.unsqueeze(dim=0).clone(), is_train=0) + ret = model.run_on_batch(rec, None) + y_pred = ret['imputations'].long().clamp(0, data.y.max()).squeeze(dim=0).T.clone() + return y_pred.clone() + + +def main(argv=None): + args = get_args(argv) + + # Match the other runners (e.g., hermes.py): run one dataset specified by CLI. + seed_all(args.seed) + tester = Tester(args.data_dir, args.device, lambda data: brits_run(data, args)) + tester.test([args.dataset], seed=args.seed, rep=1) + tester.save(args.output) + + +if __name__ == '__main__': + main() + diff --git a/brits_ms.py b/brits_ms.py new file mode 100644 index 0000000..65bb174 --- /dev/null +++ b/brits_ms.py @@ -0,0 +1,537 @@ + + +from __future__ import annotations + +import argparse +import math +from typing import List, Optional, Sequence + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.optim as optim +from torch.autograd import Variable +from torch.nn.parameter import Parameter + +try: + # Multi-snapshot tester (recommended). + from inc.test_ms import Tester +except Exception: + # Fall back to single-snapshot tester if needed. + from inc.test import Tester + +from tqdm import trange + +# Project utilities (diffusion simulator, parameter estimator, seeding, etc.) +from inc.diffus import * + + +# ----------------------------------------------------------------------------- +# CLI +# ----------------------------------------------------------------------------- + +def get_args(argv: Optional[Sequence[str]] = None) -> argparse.Namespace: + parser = argparse.ArgumentParser(description='BRITS baseline (multi-snapshot)') + + # ---- experiment I/O ---- + parser.add_argument('--dataset', type=str, required=True, help='dataset name') + parser.add_argument('--seed', type=int, default=123456789, help='random seed') + parser.add_argument('--data_dir', type=str, default='input', help='dataset folder') + parser.add_argument('--output', type=str, default='output/brits.pt', help='output file name') + parser.add_argument( + '--device', + type=torch.device, + default=torch.device('cuda' if torch.cuda.is_available() else 'cpu'), + help='torch device, e.g., cpu, cuda, cuda:0' + ) + + # ---- multi-snapshot observed times ---- + parser.add_argument( + '--obs_time', '--snapshot', + dest='obs_time', + type=str, + default='', + help='extra observed snapshot times, comma-separated, e.g., 5,7,9. ' + 'Final time T is always observed.' + ) + + # ---- diffusion parameter estimation (b_*) ---- + parser.add_argument('--b_pI0', type=float, default=1e-3, + help='initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type=float, default=1e-3, + help='initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type=int, default=500, + help='optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type=float, default=3e-3, + help='learning rate in diffusion parameter estimation') + + # ---- BRITS hyperparameters ---- + parser.add_argument('--lr', type=float, default=1e-3, help='learning rate') + parser.add_argument('--epochs', type=int, default=1000, help='training epochs') + parser.add_argument('--batch_size', type=int, default=64, help='batch size') + parser.add_argument('--hid_size', type=int, default=108, help='RNN hidden size') + parser.add_argument('--impute_weight', type=float, default=0.3, help='imputation loss weight') + parser.add_argument('--label_weight', type=float, default=1.0, help='label loss weight') + + # ---- evaluation ---- + parser.add_argument('--rep', type=int, default=1, help='number of test repetitions') + + if argv is None: + return parser.parse_args() + return parser.parse_args(argv) + + +def _parse_obs_time_str(s: str) -> List[int]: + """Parse comma-separated time indices. + + Empty string -> []. Whitespace is ignored. + """ + s = '' if s is None else str(s) + out: List[int] = [] + for tok in s.split(','): + tok = tok.strip() + if not tok: + continue + out.append(int(tok)) + return out + + +def get_obs_time(data, args: argparse.Namespace) -> List[int]: + """Return sorted unique observed times (always includes final time T).""" + T = int(data.T.item()) + obs_time = _parse_obs_time_str(getattr(args, 'obs_time', '')) + obs_time.append(T) + obs_time = sorted({t for t in obs_time if 0 <= int(t) <= T}) + if len(obs_time) == 0 or obs_time[-1] != T: + obs_time.append(T) + return obs_time + + +@torch.no_grad() +def make_obs_mask(T: int, obs_time: Sequence[int], device: torch.device) -> torch.Tensor: + """Make boolean mask of shape (T+1,) indicating observed time steps.""" + mask = torch.zeros(T + 1, dtype=torch.bool, device=device) + if len(obs_time) == 0: + mask[T] = True + return mask + ts = torch.as_tensor(list(obs_time), dtype=torch.long, device=device).clamp(0, T) + ts = ts.unique() + mask[ts] = True + if not bool(mask[T].item()): + mask[T] = True + return mask + + +# ----------------------------------------------------------------------------- +# BRITS core (largely copied from https://github.com/caow13/BRITS) +# ----------------------------------------------------------------------------- + +def binary_cross_entropy_with_logits( + input: torch.Tensor, + target: torch.Tensor, + weight: Optional[torch.Tensor] = None, + size_average: bool = True, + reduce: bool = True, +) -> torch.Tensor: + """A numerically-stable BCE-with-logits implementation. + + Kept to match the original notebook/code. + """ + if target.size() != input.size(): + raise ValueError(f'Target size ({target.size()}) must be the same as input size ({input.size()})') + max_val = (-input).clamp(min=0) + loss = input - input * target + max_val + ((-max_val).exp() + (-input - max_val).exp()).log() + if weight is not None: + loss = loss * weight + if not reduce: + return loss + if size_average: + return loss.mean() + return loss.sum() + + +class FeatureRegression(nn.Module): + def __init__(self, input_size: int): + super().__init__() + self.W = Parameter(torch.Tensor(input_size, input_size)) + self.b = Parameter(torch.Tensor(input_size)) + m = torch.ones(input_size, input_size) - torch.eye(input_size, input_size) + self.register_buffer('m', m) + self.reset_parameters() + + def reset_parameters(self) -> None: + stdv = 1.0 / math.sqrt(self.W.size(0)) + self.W.data.uniform_(-stdv, stdv) + if self.b is not None: + self.b.data.uniform_(-stdv, stdv) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + # Mask diagonal to avoid trivial self-regression + z_h = F.linear(x, self.W * Variable(self.m), self.b) + return z_h + + +class TemporalDecay(nn.Module): + def __init__(self, input_size: int, output_size: int, diag: bool = False): + super().__init__() + self.diag = diag + self.W = Parameter(torch.Tensor(output_size, input_size)) + self.b = Parameter(torch.Tensor(output_size)) + if self.diag: + assert input_size == output_size + m = torch.eye(input_size, input_size) + self.register_buffer('m', m) + self.reset_parameters() + + def reset_parameters(self) -> None: + stdv = 1.0 / math.sqrt(self.W.size(0)) + self.W.data.uniform_(-stdv, stdv) + if self.b is not None: + self.b.data.uniform_(-stdv, stdv) + + def forward(self, d: torch.Tensor) -> torch.Tensor: + if self.diag: + gamma = F.relu(F.linear(d, self.W * Variable(self.m), self.b)) + else: + gamma = F.relu(F.linear(d, self.W, self.b)) + gamma = torch.exp(-gamma) + return gamma + + +class RITS(nn.Module): + def __init__(self, xdim: int, rnn_hid_size: int, impute_weight: float, label_weight: float): + super().__init__() + self.xdim = int(xdim) + self.rnn_hid_size = int(rnn_hid_size) + self.impute_weight = float(impute_weight) + self.label_weight = float(label_weight) + + self.rnn_cell = nn.LSTMCell(self.xdim * 2, self.rnn_hid_size) + self.temp_decay_h = TemporalDecay(input_size=self.xdim, output_size=self.rnn_hid_size, diag=False) + self.temp_decay_x = TemporalDecay(input_size=self.xdim, output_size=self.xdim, diag=True) + self.hist_reg = nn.Linear(self.rnn_hid_size, self.xdim) + self.feat_reg = FeatureRegression(self.xdim) + self.weight_combine = nn.Linear(self.xdim * 2, self.xdim) + self.dropout = nn.Dropout(p=0.25) + self.out = nn.Linear(self.rnn_hid_size, 1) + + def forward(self, data: dict, direct: str): + values = data[direct]['values'] + masks = data[direct]['masks'] + deltas = data[direct]['deltas'] + evals = data[direct]['evals'] + eval_masks = data[direct]['eval_masks'] + + labels = data['labels'].reshape((-1, 1)) + is_train = data['is_train'].reshape((-1, 1)) + + device = values.device + h = Variable(torch.zeros((values.size(0), self.rnn_hid_size), device=device)) + c = Variable(torch.zeros((values.size(0), self.rnn_hid_size), device=device)) + + x_loss = 0.0 + imputations = [] + + T_seq = min(values.size(1), masks.size(1), deltas.size(1)) + for t in range(T_seq): + x = values[:, t, :] + m = masks[:, t, :] + d = deltas[:, t, :] + + gamma_h = self.temp_decay_h(d) + gamma_x = self.temp_decay_x(d) + + h = h * gamma_h + x_h = self.hist_reg(h) + + x_loss += torch.sum(torch.abs(x - x_h) * m) / (torch.sum(m) + 1e-5) + + x_c = m * x + (1 - m) * x_h + z_h = self.feat_reg(x_c) + x_loss += torch.sum(torch.abs(x - z_h) * m) / (torch.sum(m) + 1e-5) + + alpha = self.weight_combine(torch.cat([gamma_x, m], dim=1)) + c_h = alpha * z_h + (1 - alpha) * x_h + x_loss += torch.sum(torch.abs(x - c_h) * m) / (torch.sum(m) + 1e-5) + + c_c = m * x + (1 - m) * c_h + inputs = torch.cat([c_c, m], dim=1) + h, c = self.rnn_cell(inputs, (h, c)) + + imputations.append(c_c.unsqueeze(dim=1)) + + imputations = torch.cat(imputations, dim=1) + + # (unused in our setting, but keep original structure) + y_h = self.out(h) + y_loss = binary_cross_entropy_with_logits(y_h, labels, reduce=False) + y_loss = torch.sum(y_loss * is_train) / (torch.sum(is_train) + 1e-5) + y_h = torch.sigmoid(y_h) + + loss = x_loss * self.impute_weight + y_loss * self.label_weight + return { + 'loss': loss, + 'predictions': y_h, + 'imputations': imputations, + 'labels': labels, + 'is_train': is_train, + 'evals': evals, + 'eval_masks': eval_masks, + } + + def run_on_batch(self, data: dict, optimizer: Optional[optim.Optimizer], epoch: Optional[int] = None): + ret = self(data, direct='forward') + if optimizer is not None: + optimizer.zero_grad() + ret['loss'].backward() + optimizer.step() + return ret + + +class BRITSModel(nn.Module): + def __init__(self, input_size: int, hidden_size: int, impute_weight: float, label_weight: float): + super().__init__() + self.rits_f = RITS(input_size, hidden_size, impute_weight, label_weight) + self.rits_b = RITS(input_size, hidden_size, impute_weight, label_weight) + + def forward(self, data: dict): + ret_f = self.rits_f(data, direct='forward') + ret_b = self.rits_b(data, direct='backward') + ret_b = self.reverse(ret_b) + + loss_f = ret_f['loss'] + loss_b = ret_b['loss'] + loss_c = self.get_consistency_loss(ret_f['imputations'], ret_b['imputations']) + loss = loss_f + loss_b + loss_c + + predictions = (ret_f['predictions'] + ret_b['predictions']) / 2 + imputations = (ret_f['imputations'] + ret_b['imputations']) / 2 + + ret_f['loss'] = loss + ret_f['predictions'] = predictions + ret_f['imputations'] = imputations + return ret_f + + @staticmethod + def get_consistency_loss(pred_f: torch.Tensor, pred_b: torch.Tensor) -> torch.Tensor: + return torch.abs(pred_f - pred_b).mean() * 1e-1 + + @staticmethod + def reverse(ret: dict) -> dict: + def reverse_tensor(tensor_: torch.Tensor) -> torch.Tensor: + if tensor_.dim() <= 1: + return tensor_ + # reverse along time dimension (dim=1) + idx = torch.arange(tensor_.size(1) - 1, -1, -1, device=tensor_.device) + return tensor_.index_select(1, idx) + + return {k: reverse_tensor(v) for k, v in ret.items()} + + def run_on_batch(self, data: dict, optimizer: Optional[optim.Optimizer], epoch: Optional[int] = None): + ret = self(data) + if optimizer is not None: + optimizer.zero_grad() + ret['loss'].backward() + optimizer.step() + return ret + + +class BRITS(nn.Module): + def __init__(self, input_size: int, hidden_size: int, impute_weight: float = 1.0, label_weight: float = 1.0): + super().__init__() + self.model = BRITSModel(input_size, hidden_size, impute_weight, label_weight) + + def forward(self, data: dict): + return self.model(data) + + def run_on_batch(self, data: dict, optimizer: Optional[optim.Optimizer], epoch: Optional[int] = None): + return self.model.run_on_batch(data, optimizer, epoch) + + +# ----------------------------------------------------------------------------- +# Data preparation (multi-snapshot masking) +# ----------------------------------------------------------------------------- + +def to_var(var): + """Recursively move nested structures into torch Variables on the same device.""" + if torch.is_tensor(var): + dev = var.device + var = torch.autograd.Variable(var) + var = var.to(dev) + return var + if isinstance(var, (int, float, str)): + return var + if isinstance(var, dict): + return {key: to_var(val) for key, val in var.items()} + if isinstance(var, list): + return [to_var(x) for x in var] + return var + + +@torch.no_grad() +def _brits_make_deltas(obs_mask_1d: torch.Tensor) -> torch.Tensor: + """Make a (T+1,) float tensor of time gaps since the last observation. + + In forward direction: + delta[t] = 0 if t is observed else delta[t-1] + 1. + In backward direction we apply the same logic on the reversed mask. + """ + if obs_mask_1d.dim() != 1: + raise ValueError(f'obs_mask_1d must be 1D, got shape={tuple(obs_mask_1d.shape)}') + + L = obs_mask_1d.numel() + d = torch.zeros(L, dtype=torch.float, device=obs_mask_1d.device) + for t in range(1, L): + d[t] = 0.0 if bool(obs_mask_1d[t].item()) else (d[t - 1] + 1.0) + return d + + +@torch.no_grad() +def brits_prep_rec(y: torch.Tensor, obs_mask_1d: torch.Tensor) -> dict: + """Prepare one direction (forward OR backward) input for BRITS. + + Args: + y: (samples, T+1, nodes) full ground-truth sequence (will be masked). + obs_mask_1d: (T+1,) bool mask for which time steps are observed. + + Returns: + dict with keys values/masks/evals/eval_masks/deltas, all shaped + (samples, T+1, nodes). + """ + if y.dim() != 3: + raise ValueError(f'y must be 3D (samples,T+1,nodes), got {tuple(y.shape)}') + + # Eval targets (full sequence, used only for reporting; NOT fed as observed). + evals = y.float().clone() + + # Observation mask broadcast to (samples, T+1, nodes) + obs_mask = obs_mask_1d.to(device=y.device, dtype=torch.bool) + masks = obs_mask.view(1, -1, 1).expand(y.size(0), -1, y.size(2)).float() + + # Only keep observed frames in values; unobserved are set to 0. + values = evals * masks + + # delta features + deltas_1d = _brits_make_deltas(obs_mask) + deltas = deltas_1d.view(1, -1, 1).expand_as(values).contiguous() + + eval_masks = masks.clone() + + return { + 'values': values.contiguous(), + 'masks': masks.contiguous(), + 'evals': evals.contiguous(), + 'eval_masks': eval_masks.contiguous(), + 'deltas': deltas.contiguous(), + } + + +@torch.no_grad() +def brits_prep(y: torch.Tensor, is_train: int, obs_mask_1d: torch.Tensor) -> dict: + """Prepare BRITS batch dict. + + Args: + y: (samples, nodes, T+1) + is_train: 1 or 0 + obs_mask_1d: (T+1,) bool in *forward* time. + + Returns: + dict consumable by BRITSModel. + """ + if y.dim() != 3: + raise ValueError(f'y must be 3D (samples,nodes,T+1), got {tuple(y.shape)}') + + n_samples = int(y.size(0)) + y = y.transpose(1, 2) # (samples, T+1, nodes) + + obs_mask_1d = obs_mask_1d.to(device=y.device, dtype=torch.bool) + obs_mask_bwd = obs_mask_1d.flip(dims=[0]) + + return to_var({ + 'forward': brits_prep_rec(y, obs_mask_1d), + 'backward': brits_prep_rec(y.flip(dims=[1]), obs_mask_bwd), + # BRITS classification head is unused; keep placeholders + 'labels': torch.zeros(n_samples, 1, dtype=torch.long, device=y.device), + 'is_train': torch.full((n_samples, 1), float(is_train), dtype=torch.float, device=y.device), + }) + + +# ----------------------------------------------------------------------------- +# Experiment runner +# ----------------------------------------------------------------------------- + +def brits_run(data, args: argparse.Namespace) -> torch.Tensor: + """Train BRITS on synthetic histories and infer a history for the given data. + + Returns: + y_pred: (nodes, T+1) long tensor of reconstructed states. + """ + device = args.device + data = data.to(device) + + T = int(data.T.item()) + obs_time = get_obs_time(data, args) + + # Attach observation times for the tester (so metrics exclude observed frames) + # `inc.test_ms` uses `data.obs_ts` or `data.obs_mask` if present. + data.obs_ts = obs_time + + obs_mask = make_obs_mask(T, obs_time, device=device) + + # Estimate diffusion parameters (supports multi-snapshot via obs_time) + bpar = b_estim(data, args, obs_time=obs_time) + + model = BRITS(data.num_nodes, args.hid_size, args.impute_weight, args.label_weight).to(device) + + # -------------------- train -------------------- + I0 = int((data.y[:, 0] == SIR_STATES.I).long().sum().item()) + optimizer = optim.Adam(model.parameters(), lr=args.lr) + + pbar = trange(args.epochs) + for epoch in pbar: + model.train() + # Generate synthetic training batch (full history), then mask it. + Y = diffus_gen( + T=T, + n_nodes=data.num_nodes, + edge_index=data.edge_index, + I0=I0, + n_samples=args.batch_size, + pI=bpar.pI, + pR=bpar.pR, + ) # (T+1, nodes, samples) + + batch_y = Y.transpose(0, 2).clone() # -> (samples, nodes, T+1) + batch = brits_prep(batch_y, is_train=1, obs_mask_1d=obs_mask) + ret = model.run_on_batch(batch, optimizer, epoch) + pbar.set_description(f'epoch={epoch + 1} loss={ret["loss"].item():.4f}') + + # -------------------- infer -------------------- + with torch.no_grad(): + model.eval() + rec = brits_prep(data.y.unsqueeze(dim=0).clone(), is_train=0, obs_mask_1d=obs_mask) + ret = model.run_on_batch(rec, None) + + # ret['imputations']: (1, T+1, nodes) float + y_pred = ret['imputations'] + y_pred = y_pred.long().clamp(0, int(data.y.max().item())) + y_pred = y_pred.squeeze(dim=0).T.contiguous() # (nodes, T+1) + return y_pred.detach().clone() + + +def main(argv: Optional[Sequence[str]] = None) -> None: + args = get_args(argv) + + # Reproducibility + seed_all(args.seed) + + def _model_fn(data): + return brits_run(data, args) + + tester = Tester(args.data_dir, args.device, _model_fn) + tester.test([args.dataset], seed=args.seed, rep=args.rep) + tester.save(args.output) + + +if __name__ == '__main__': + main() diff --git a/grin.py b/grin.py new file mode 100644 index 0000000..ed90ba4 --- /dev/null +++ b/grin.py @@ -0,0 +1,219 @@ +# -*- coding: utf-8 -*- +# NOTE: This file was extracted from `grin.ipynb`. +# - Spatiotemporal 0.1.1: https://github.com/TorchSpatiotemporal/tsl/tree/1ae3289e00b28d0e84dfd54799561162df1917cd +# - SPIN: https://github.com/Graph-Machine-Learning-Group/spin + +from __future__ import annotations + +import argparse + +import torch +from tsl.nn.models.stgn import GRINModel + +from inc.diffus import * +from inc.test import * + + +def get_args() -> argparse.Namespace: + """Parse command line arguments.""" + + parser = argparse.ArgumentParser(description='GRIN baseline (single-snapshot)') + + # ---- standard experiment args (align with other runners, e.g., hermes.py) ---- + parser.add_argument('--dataset', type=str, required=True, help='dataset name') + parser.add_argument('--seed', type=int, default=123456789, help='random seed') + parser.add_argument('--data_dir', type=str, default='input', help='dataset folder') + parser.add_argument('--output', type=str, default='output/grin.pt', help='output file name') + parser.add_argument( + '--device', + type=str, + default='cuda' if torch.cuda.is_available() else 'cpu', + help='torch device, e.g., cpu | cuda | cuda:0', + ) + + # ---- diffusion parameter estimation (b_*) ---- + parser.add_argument( + '--b_pI0', + type=float, + default=1e-3, + help='initial infection rate in diffusion parameter estimation', + ) + parser.add_argument( + '--b_pR0', + type=float, + default=1e-3, + help='initial recovery rate in diffusion parameter estimation', + ) + parser.add_argument( + '--b_steps', + type=int, + default=500, + help='optimization steps in diffusion parameter estimation', + ) + parser.add_argument( + '--b_lr', + type=float, + default=3e-3, + help='learning rate in diffusion parameter estimation', + ) + + # ---- GRINModel hyper-parameters ---- + # Ref: https://github.com/Graph-Machine-Learning-Group/spin/blob/main/config/imputation/grin.yaml + parser.add_argument('--hidden_size', type=int, default=64) + parser.add_argument('--ff_size', type=int, default=64) + parser.add_argument('--embedding_size', type=int, default=8) + parser.add_argument('--n_layers', type=int, default=1) + parser.add_argument('--kernel_size', type=int, default=2) + parser.add_argument('--decoder_order', type=int, default=1) + parser.add_argument( + '--layer_norm', + action='store_true', + help='enable layer norm inside the GRIN model', + ) + parser.add_argument('--dropout', type=float, default=0.0) + parser.add_argument('--ff_dropout', type=float, default=0.0) + parser.add_argument('--merge_mode', type=str, default='mlp') + + # ---- optimizer / training ---- + parser.add_argument('--lr', type=float, default=1e-3, help='Adam learning rate') + parser.add_argument('--l2_reg', type=float, default=0.0, help='Adam weight decay') + parser.add_argument('--epochs', type=int, default=300) + parser.add_argument('--batch_size', type=int, default=1) + + # ---- evaluation ---- + parser.add_argument('--rep', type=int, default=1, help='repetitions per dataset') + + args = parser.parse_args() + + # Normalize device to torch.device (keeps compatibility with other modules). + args.device = torch.device(args.device) + + return args + + +def grin_prep(y: torch.Tensor, edge_index: torch.Tensor, device: torch.device): + """Prepare inputs for GRIN. + + Args: + y: (samples, nodes, T+1) integer states. + edge_index: (2, edges) + device: torch device for mask tensor + + Returns: + x: (samples, T+1, nodes, 1) + mask: (samples, T+1, nodes, 1) with only last snapshot observed + ei: edge_index (kept for API symmetry) + """ + + n_samples, n_nodes, T1 = y.size() + T = T1 - 1 + + # (samples, T+1, nodes, 1) + x = y.transpose(1, 2).unsqueeze(dim=3).float() + + # only last snapshot observed + mask = ( + torch.tensor([[0]] * T + [[1]], device=device) + .expand(n_samples, n_nodes, -1, 1) + .transpose(1, 2) + ) + + ei = edge_index + return x, mask, ei + + +def grin_run(data, args: argparse.Namespace): + """Train GRIN on synthetic histories generated from estimated diffusion params, + then impute the unobserved history for `data`. + + This is a minimal refactor of the original notebook logic. + """ + + bpar = b_estim(data, args) # Dict(pI=..., pR=...) + + T = data.T.item() + n_nodes = data.num_nodes + n_out = data.y[:, -1].max().item() + 1 + + # train + model = GRINModel( + input_size=1, + hidden_size=args.hidden_size, + ff_size=args.ff_size, + embedding_size=args.embedding_size, + n_layers=args.n_layers, + n_nodes=n_nodes, + kernel_size=args.kernel_size, + decoder_order=args.decoder_order, + layer_norm=args.layer_norm, + dropout=args.dropout, + ff_dropout=args.ff_dropout, + merge_mode=args.merge_mode, + ).to(args.device) + + I0 = (data.y[:, 0] == SIR_STATES.I).long().sum().item() + + opt = torch.optim.Adam(model.parameters(), lr=args.lr, weight_decay=args.l2_reg) + model.train() + + pbar = trange(1, args.epochs + 1) + for epoch in pbar: + opt.zero_grad() + + # generate synthetic histories for training + Y_true = diffus_gen( + T=data.T.item(), + n_nodes=data.num_nodes, + edge_index=data.edge_index, + I0=I0, + n_samples=args.batch_size, + pI=bpar.pI, + pR=bpar.pR, + ) # (T+1, nodes, samples) + + # GRIN expects (samples, T+1, nodes, features) + x, mask, _ = grin_prep(Y_true.transpose(0, 2), data.edge_index, device=args.device) + + z = model(x=x, mask=mask, edge_index=data.edge_index)[0] # (samples, T+1, nodes, 1) + + # L1 loss on the unobserved part (t < T) + loss = ( + z[:, :-1].flatten() + - Y_true[:-1].transpose(1, 2).transpose(0, 1).flatten() + ).abs().mean() + + pbar.set_description(f'[epoch={epoch}] loss={loss.item():.4f}') + loss.backward() + opt.step() + + # infer + with torch.no_grad(): + model.eval() + x, mask, ei = grin_prep(data.y.unsqueeze(dim=0).clone(), data.edge_index, device=args.device) + z = model(x=x, mask=mask, edge_index=ei)[0] # (1, T+1, nodes, 1) + + y_pred = data.y.clone() + y_pred[:, :-1] = ( + z[0, :-1, :, 0] + .clamp(0, n_out - 1) + .T + .round() + .long() + ) # (nodes, T) + + return y_pred.clone() + + +def main() -> None: + args = get_args() + + # Tester expects a callable: model_fn(data) -> y_pred + model_fn = lambda data: grin_run(data, args) + + tester = Tester(args.data_dir, args.device, model_fn) + tester.test([args.dataset], seed=args.seed, rep=args.rep) + tester.save(args.output) + + +if __name__ == '__main__': + main() diff --git a/grin_ms.py b/grin_ms.py new file mode 100644 index 0000000..efb624a --- /dev/null +++ b/grin_ms.py @@ -0,0 +1,272 @@ + +from __future__ import annotations + +import argparse +import os +from typing import List + +import torch +from tqdm import trange +from tsl.nn.models.stgn import GRINModel + +# Project utilities (same style as other runners) +from inc.diffus import SIR_STATES, b_estim, diffus_gen, seed_all + +# Multi-snapshot aware tester (only evaluates on unobserved positions) +# If your repo uses inc.test as the ms tester, you can swap this import accordingly. +from inc.test_ms import Tester + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description="GRIN baseline (multi-snapshot)") + + # ---- standard experiment args (align with hermes.py style) ---- + parser.add_argument("--dataset", type=str, required=True, help="dataset name") + parser.add_argument("--seed", type=int, default=123456789, help="random seed") + parser.add_argument("--data_dir", type=str, default="input", help="dataset folder") + parser.add_argument("--output", type=str, default="output/grin.pt", help="output file name") + parser.add_argument( + "--device", + type=str, + default="cuda" if torch.cuda.is_available() else "cpu", + help="torch device, e.g., cpu | cuda | cuda:0", + ) + + # ---- multi-snapshot control ---- + # Keep hermes-style naming (obs_time) and also provide snapshot as an alias. + parser.add_argument( + "--obs_time", + "--snapshot", + dest="obs_time", + type=str, + default="", + help="extra observed snapshot times, comma-separated (e.g., 5 or 3,7,9). " + "Final time T is always observed automatically.", + ) + + # ---- diffusion parameter estimation (b_*) ---- + parser.add_argument("--b_pI0", type=float, default=1e-3, + help="initial infection rate in diffusion parameter estimation") + parser.add_argument("--b_pR0", type=float, default=1e-3, + help="initial recovery rate in diffusion parameter estimation") + parser.add_argument("--b_steps", type=int, default=500, + help="optimization steps in diffusion parameter estimation") + parser.add_argument("--b_lr", type=float, default=3e-3, + help="learning rate in diffusion parameter estimation") + + # ---- GRINModel hyper-parameters ---- + # Ref: SPIN config/imputation/grin.yaml + parser.add_argument("--hidden_size", type=int, default=64) + parser.add_argument("--ff_size", type=int, default=64) + parser.add_argument("--embedding_size", type=int, default=8) + parser.add_argument("--n_layers", type=int, default=1) + parser.add_argument("--kernel_size", type=int, default=2) + parser.add_argument("--decoder_order", type=int, default=1) + parser.add_argument("--layer_norm", action="store_true", + help="enable layer norm inside GRIN") + parser.add_argument("--dropout", type=float, default=0.0) + parser.add_argument("--ff_dropout", type=float, default=0.0) + parser.add_argument("--merge_mode", type=str, default="mlp") + + # ---- optimizer / training ---- + parser.add_argument("--lr", type=float, default=1e-3, help="Adam learning rate") + parser.add_argument("--l2_reg", type=float, default=0.0, help="Adam weight decay") + parser.add_argument("--epochs", type=int, default=300) + parser.add_argument("--batch_size", type=int, default=1) + + # ---- evaluation ---- + parser.add_argument("--rep", type=int, default=1, help="repetitions per dataset") + + return parser + + +def parse_obs_time(obs_time_str: str, T: int) -> List[int]: + """ + Parse comma-separated times from CLI, clamp to [0, T], and always include T. + + hermes.py behavior: + obs_time = [int(t) for t in args.obs_time.split(',') if t] + obs_time.append(T) + obs_time = sorted(set(obs_time)) + """ + times: List[int] = [] + if obs_time_str: + for s in str(obs_time_str).split(","): + s = s.strip() + if not s: + continue + try: + t = int(s) + except ValueError: + continue + if 0 <= t <= T: + times.append(t) + + times.append(T) # final snapshot always observed + times = sorted(set(times)) + return times + + +def build_time_mask(obs_time: List[int], T: int, device: torch.device) -> torch.Tensor: + """ + Return a 1D mask over time: (T+1,) with 1 at observed times, else 0. + """ + m = torch.zeros(T + 1, dtype=torch.long, device=device) + if len(obs_time) > 0: + idx = torch.tensor(obs_time, dtype=torch.long, device=device).clamp(0, T) + m[idx.unique()] = 1 + return m + + +def grin_prep(y: torch.Tensor, obs_time: List[int], device: torch.device): + """ + Prepare inputs for GRIN. + + Args: + y: (samples, nodes, T+1) integer states. + obs_time: list of observed snapshot times (must include T). + device: torch device for mask. + + Returns: + x: (samples, T+1, nodes, 1) float, with missing frames zeroed + mask: (samples, T+1, nodes, 1) long {0,1}, 1 means observed + """ + n_samples, n_nodes, T1 = y.size() + T = T1 - 1 + + # time mask: (T+1,) + tmask = build_time_mask(obs_time, T=T, device=device) # long (T+1,) + + # GRIN input layout: (samples, T+1, nodes, 1) + x = y.transpose(1, 2).unsqueeze(dim=3).float() + + # mask layout: (samples, T+1, nodes, 1) + mask = tmask.view(1, T1, 1, 1).expand(n_samples, T1, n_nodes, 1) + + # IMPORTANT: avoid leaking ground-truth values at missing times + x = x * mask.float() + + return x, mask + + +def grin_run_ms(data, args: argparse.Namespace) -> torch.Tensor: + """ + Train GRIN on synthetic histories generated from estimated diffusion params, + then impute the unobserved history for `data` under the multi-snapshot mask. + """ + # ---- parse observed snapshots ---- + T = int(data.T.item()) + obs_time = parse_obs_time(args.obs_time, T) + + # attach obs info for the ms tester (so metrics only evaluate unobserved frames) + # test_ms.py will look for data.obs_ts / obs_mask / obs_masks + data.obs_ts = obs_time + + # ---- estimate diffusion params using observed snapshots ---- + bpar = b_estim(data, args, obs_time=obs_time) # Dict(pI=..., pR=...) + + n_nodes = int(data.num_nodes) + n_out = int(data.y[:, -1].max().item() + 1) + + # ---- build model ---- + model = GRINModel( + input_size=1, + hidden_size=args.hidden_size, + ff_size=args.ff_size, + embedding_size=args.embedding_size, + n_layers=args.n_layers, + n_nodes=n_nodes, + kernel_size=args.kernel_size, + decoder_order=args.decoder_order, + layer_norm=args.layer_norm, + dropout=args.dropout, + ff_dropout=args.ff_dropout, + merge_mode=args.merge_mode, + ).to(args.device) + + # prior info used in original baseline (kept) + I0 = int((data.y[:, 0] == SIR_STATES.I).long().sum().item()) + + opt = torch.optim.Adam(model.parameters(), lr=args.lr, weight_decay=args.l2_reg) + model.train() + + # ---- training ---- + pbar = trange(1, args.epochs + 1) + for epoch in pbar: + opt.zero_grad() + + # synthetic training histories: (T+1, nodes, samples) + Y_true = diffus_gen( + T=T, + n_nodes=n_nodes, + edge_index=data.edge_index, + I0=I0, + n_samples=args.batch_size, + pI=bpar.pI, + pR=bpar.pR, + ) + + # GRIN expects (samples, T+1, nodes, 1) + # y for prep: (samples, nodes, T+1) + y_snt = Y_true.transpose(0, 2) # (samples, nodes, T+1) + x, mask = grin_prep(y_snt, obs_time=obs_time, device=args.device) + + # model output: (samples, T+1, nodes, 1) + z = model(x=x, mask=mask, edge_index=data.edge_index)[0] + + # ground-truth in same layout: (samples, T+1, nodes, 1) + y_true_seq = Y_true.permute(2, 0, 1).unsqueeze(-1).float() + + # loss only on missing frames (mask == 0) + unobs = (mask == 0) + if bool(unobs.any()): + loss = (z - y_true_seq).abs()[unobs].mean() + else: + # degenerate case: everything observed (rare), fallback to full loss + loss = (z - y_true_seq).abs().mean() + + pbar.set_description(f"[epoch={epoch}] loss={loss.item():.4f}") + loss.backward() + opt.step() + + # ---- inference ---- + with torch.no_grad(): + model.eval() + + # build masked input from observed snapshots + tmask_1d = build_time_mask(obs_time, T=T, device=args.device).bool() # (T+1,) + y_in = data.y.clone() + y_in[:, ~tmask_1d] = 0 # hide unobserved frames + + x, mask = grin_prep(y_in.unsqueeze(0), obs_time=obs_time, device=args.device) + + z = model(x=x, mask=mask, edge_index=data.edge_index)[0] # (1, T+1, nodes, 1) + + pred = z[0, :, :, 0].T # (nodes, T+1) + pred = pred.clamp(0, n_out - 1).round().long() + + y_pred = data.y.clone() + y_pred[:, ~tmask_1d] = pred[:, ~tmask_1d] # only fill missing times + return y_pred + + +def main() -> None: + parser = build_parser() + args = parser.parse_args() + + # normalize device + args.device = torch.device(args.device) + + # make sure output dir exists + out_dir = os.path.dirname(args.output) + if out_dir: + os.makedirs(out_dir, exist_ok=True) + + # run + tester = Tester(args.data_dir, args.device, lambda data: grin_run_ms(data, args)) + tester.test([args.dataset], seed=args.seed, rep=args.rep) + tester.save(args.output) + + +if __name__ == "__main__": + main() diff --git a/spin.py b/spin.py new file mode 100644 index 0000000..0fd251d --- /dev/null +++ b/spin.py @@ -0,0 +1,672 @@ +# ! pip install --no-index torch-scatter==2.0.7 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +# ! pip install --no-index torch-sparse==0.6.9 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +# ! pip install --no-index torch-cluster==1.5.9 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +#! pip install --no-index torch-spline-conv==1.2.1 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +# ! pip install torch-geometric==2.0.4 +# ! pip install ndlib==5.1.1 + + +# ! pip install einops==0.6.0 +# ! pip install test_tube==0.7.5 + + +try: + from tsl.nn.base import StaticGraphEmbedding +except Exception: + import torch + from torch import nn + + class StaticGraphEmbedding(nn.Module): + def __init__(self, n_nodes, out_channels): + super().__init__() + self.emb = nn.Embedding(n_nodes, out_channels) + + def forward(self, token_index=None): + if token_index is None: + token_index = torch.arange(self.emb.num_embeddings, device=self.emb.weight.device) + return self.emb(token_index) +from tsl.nn.layers import PositionalEncoding +from tsl.nn.layers.norm import LayerNorm +from tsl.nn.blocks.encoders import MLP +from tsl.nn.functional import sparse_softmax +from tsl.engines import Imputer, Predictor +from tsl.ops.connectivity import weighted_degree +#from tsl.data import Batch, SpatioTemporalDataModule, ImputationDataset + +#SPINModel +'''https://github.com/Graph-Machine-Learning-Group/spin/blob/main/spin/layers/postional_encoding.py''' +from typing import Optional + +from torch import nn + +class PositionalEncoder(nn.Module): + + def __init__(self, in_channels, out_channels, + n_layers: int = 1, + n_nodes: Optional[int] = None): + super(PositionalEncoder, self).__init__() + self.lin = nn.Linear(in_channels, out_channels) + self.activation = nn.LeakyReLU() + self.mlp = MLP(out_channels, out_channels, out_channels, + n_layers=n_layers, activation='relu') + self.positional = PositionalEncoding(out_channels) + if n_nodes is not None: + self.node_emb = StaticGraphEmbedding(n_nodes, out_channels) + else: + self.register_parameter('node_emb', None) + + def forward(self, x, node_emb=None, node_index=None): + if node_emb is None: + node_emb = self.node_emb(token_index=node_index) + # x: [b s c], node_emb: [n c] -> [b s n c] + x = self.lin(x) + x = self.activation(x.unsqueeze(-2) + node_emb) + #print('u:', tuple(x.shape), 'node_emb:', tuple(node_emb.shape))##### + out = self.mlp(x) + out = self.positional(out) + return out + +'''https://github.com/Graph-Machine-Learning-Group/spin/blob/main/spin/layers/additive_attention.py''' +from typing import Optional, Tuple, Union + +import torch +from torch import Tensor +from torch import nn +from torch.nn import LayerNorm, functional as F +from torch_geometric.nn.conv import MessagePassing +from torch_geometric.nn.dense.linear import Linear +from torch_geometric.typing import Adj, OptTensor, PairTensor +from torch_scatter import scatter +from torch_scatter.utils import broadcast + + +class AdditiveAttention(MessagePassing): + def __init__(self, input_size: Union[int, Tuple[int, int]], + output_size: int, + msg_size: Optional[int] = None, + msg_layers: int = 1, + root_weight: bool = True, + reweight: Optional[str] = None, + norm: bool = True, + dropout: float = 0.0, + dim: int = -2, + **kwargs): + kwargs.setdefault('aggr', 'add') + super().__init__(node_dim=dim, **kwargs) + + self.output_size = output_size + if isinstance(input_size, int): + self.src_size = self.tgt_size = input_size + else: + self.src_size, self.tgt_size = input_size + + self.msg_size = msg_size or self.output_size + self.msg_layers = msg_layers + + assert reweight in ['softmax', 'l1', None] + self.reweight = reweight + + self.root_weight = root_weight + self.dropout = dropout + + # key bias is discarded in softmax + self.lin_src = Linear(self.src_size, self.output_size, + weight_initializer='glorot', + bias_initializer='zeros') + self.lin_tgt = Linear(self.tgt_size, self.output_size, + weight_initializer='glorot', bias=False) + + if self.root_weight: + self.lin_skip = Linear(self.tgt_size, self.output_size, + bias=False) + else: + self.register_parameter('lin_skip', None) + + self.msg_nn = nn.Sequential( + nn.PReLU(init=0.2), + MLP(self.output_size, self.msg_size, self.output_size, + n_layers=self.msg_layers, dropout=self.dropout, + activation='prelu') + ) + + if self.reweight == 'softmax': + self.msg_gate = nn.Linear(self.output_size, 1, bias=False) + else: + self.msg_gate = nn.Sequential(nn.Linear(self.output_size, 1), + nn.Sigmoid()) + + if norm: + self.norm = LayerNorm(self.output_size) + else: + self.register_parameter('norm', None) + + self.reset_parameters() + + def reset_parameters(self): + self.lin_src.reset_parameters() + self.lin_tgt.reset_parameters() + if self.lin_skip is not None: + self.lin_skip.reset_parameters() + + def forward(self, x: PairTensor, edge_index: Adj, mask: OptTensor = None): + # if query/key not provided, defaults to x (e.g., for self-attention) + if isinstance(x, Tensor): + x_src = x_tgt = x + else: + x_src, x_tgt = x + x_tgt = x_tgt if x_tgt is not None else x_src + + N_src, N_tgt = x_src.size(self.node_dim), x_tgt.size(self.node_dim) + + msg_src = self.lin_src(x_src) + msg_tgt = self.lin_tgt(x_tgt) + + msg = (msg_src, msg_tgt) + + # propagate_type: (msg: PairTensor, mask: OptTensor) + out = self.propagate(edge_index, msg=msg, mask=mask, + size=(N_src, N_tgt)) + + # skip connection + if self.root_weight: + out = out + self.lin_skip(x_tgt) + + if self.norm is not None: + out = self.norm(out) + + return out + + def normalize_weights(self, weights, index, num_nodes, mask=None): + # mask weights + if mask is not None: + fill_value = float("-inf") if self.reweight == 'softmax' else 0. + weights = weights.masked_fill(torch.logical_not(mask), fill_value) + # eventually reweight + if self.reweight == 'l1': + expanded_index = broadcast(index, weights, self.node_dim) + weights_sum = scatter(weights, expanded_index, self.node_dim, + dim_size=num_nodes, reduce='sum') + weights_sum = weights_sum.index_select(self.node_dim, index) + weights = weights / (weights_sum + 1e-5) + elif self.reweight == 'softmax': + weights = sparse_softmax(weights, index, num_nodes=num_nodes, + dim=self.node_dim) + return weights + + def message(self, msg_j: Tensor, msg_i: Tensor, index, size_i, + mask_j: OptTensor = None) -> Tensor: + msg = self.msg_nn(msg_j + msg_i) + gate = self.msg_gate(msg) + alpha = self.normalize_weights(gate, index, size_i, mask_j) + alpha = F.dropout(alpha, p=self.dropout, training=self.training) + out = alpha * msg + return out + + def __repr__(self) -> str: + return (f'{self.__class__.__name__}({self.output_size}, ' + f'dim={self.node_dim}, ' + f'root_weight={self.root_weight})') + + +class TemporalAdditiveAttention(AdditiveAttention): + def __init__(self, input_size: Union[int, Tuple[int, int]], + output_size: int, + msg_size: Optional[int] = None, + msg_layers: int = 1, + root_weight: bool = True, + reweight: Optional[str] = None, + norm: bool = True, + dropout: float = 0.0, + **kwargs): + kwargs.setdefault('dim', 1) + super().__init__(input_size=input_size, + output_size=output_size, + msg_size=msg_size, + msg_layers=msg_layers, + root_weight=root_weight, + reweight=reweight, + dropout=dropout, + norm=norm, + **kwargs) + + def forward(self, x: PairTensor, mask: OptTensor = None, + temporal_mask: OptTensor = None, + causal_lag: Optional[int] = None): + # x: [b s * c] query: [b l * c] key: [b s * c] + # mask: [b s * c] temporal_mask: [l s] + if isinstance(x, Tensor): + x_src = x_tgt = x + else: + x_src, x_tgt = x + x_tgt = x_tgt if x_tgt is not None else x_src + + l, s = x_tgt.size(self.node_dim), x_src.size(self.node_dim) + i = torch.arange(l, dtype=torch.long, device=x_src.device) + j = torch.arange(s, dtype=torch.long, device=x_src.device) + + # compute temporal index, from j to i + if temporal_mask is None and isinstance(causal_lag, int): + temporal_mask = tuple(torch.tril_indices(l, l, offset=-causal_lag, + device=x_src.device)) + if temporal_mask is not None: + assert temporal_mask.size() == (l, s) + i, j = torch.meshgrid(i, j) + edge_index = torch.stack((j[temporal_mask], i[temporal_mask])) + else: + edge_index = torch.cartesian_prod(j, i).T + + return super(TemporalAdditiveAttention, self).forward(x, edge_index, + mask=mask) + +'''https://github.com/Graph-Machine-Learning-Group/spin/blob/main/spin/layers/temporal_graph_additive_attention.py''' +from typing import Optional, Tuple, Union + +import torch +from torch import Tensor +from torch_geometric.nn.conv import MessagePassing +from torch_geometric.nn.dense.linear import Linear +from torch_geometric.typing import Adj, OptTensor, OptPairTensor + +class TemporalGraphAdditiveAttention(MessagePassing): + def __init__(self, input_size: Union[int, Tuple[int, int]], + output_size: int, + msg_size: Optional[int] = None, + msg_layers: int = 1, + root_weight: bool = True, + reweight: Optional[str] = None, + temporal_self_attention: bool = True, + mask_temporal: bool = True, + mask_spatial: bool = True, + norm: bool = True, + dropout: float = 0., + **kwargs): + kwargs.setdefault('aggr', 'add') + super(TemporalGraphAdditiveAttention, self).__init__(node_dim=-2, + **kwargs) + + # store dimensions + if isinstance(input_size, int): + self.src_size = self.tgt_size = input_size + else: + self.src_size, self.tgt_size = input_size + self.output_size = output_size + self.msg_size = msg_size or self.output_size + + self.mask_temporal = mask_temporal + self.mask_spatial = mask_spatial + + self.root_weight = root_weight + self.dropout = dropout + + if temporal_self_attention: + self.self_attention = TemporalAdditiveAttention( + input_size=input_size, + output_size=output_size, + msg_size=msg_size, + msg_layers=msg_layers, + reweight=reweight, + dropout=dropout, + root_weight=False, + norm=False + ) + else: + self.register_parameter('self_attention', None) + + self.cross_attention = TemporalAdditiveAttention(input_size=input_size, + output_size=output_size, + msg_size=msg_size, + msg_layers=msg_layers, + reweight=reweight, + dropout=dropout, + root_weight=False, + norm=False) + + if self.root_weight: + self.lin_skip = Linear(self.tgt_size, self.output_size, + bias_initializer='zeros') + else: + self.register_parameter('lin_skip', None) + + if norm: + self.norm = LayerNorm(output_size) + else: + self.register_parameter('norm', None) + + self.reset_parameters() + + def reset_parameters(self): + self.cross_attention.reset_parameters() + if self.self_attention is not None: + self.self_attention.reset_parameters() + if self.lin_skip is not None: + self.lin_skip.reset_parameters() + if self.norm is not None: + self.norm.reset_parameters() + + def forward(self, x: OptPairTensor, + edge_index: Adj, edge_weight: OptTensor = None, + mask: OptTensor = None): + # inputs: [batch, steps, nodes, channels] + if isinstance(x, Tensor): + x_src = x_tgt = x + else: + x_src, x_tgt = x + x_tgt = x_tgt if x_tgt is not None else x_src + + n_src, n_tgt = x_src.size(-2), x_tgt.size(-2) + + # propagate query, key and value + #print('src:', x_src.shape, 'tgt:', x_tgt.shape, 'ei:', edge_index.shape, 'mask:', mask.shape, f'mask_spatial={self.mask_spatial}') + out = self.propagate(x=(x_src, x_tgt), + edge_index=edge_index, edge_weight=edge_weight, + mask=mask if self.mask_spatial else None, + size=(n_src, n_tgt)) + + if self.self_attention is not None: + s, l = x_src.size(1), x_tgt.size(1) + if s == l: + attn_mask = ~torch.eye(l, l, dtype=torch.bool, + device=x_tgt.device) + else: + attn_mask = None + temp = self.self_attention(x=(x_src, x_tgt), + mask=mask if self.mask_temporal else None, + temporal_mask=attn_mask) + out = out + temp + + # skip connection + if self.root_weight: + out = out + self.lin_skip(x_tgt) + + if self.norm is not None: + out = self.norm(out) + + return out + + def message(self, x_i: Tensor, x_j: Tensor, + edge_weight: OptTensor, mask_j: OptTensor) -> Tensor: + # [batch, steps, edges, channels] + + out = self.cross_attention((x_j, x_i), mask=mask_j) + #print('out:', out.shape) + + if edge_weight is not None: + out = out * edge_weight.view(-1, 1) + return out + +'''https://github.com/Graph-Machine-Learning-Group/spin/blob/main/spin/models/spin.py''' +from typing import Optional + +import torch +from torch import nn, Tensor +from torch.nn import LayerNorm +from torch_geometric.typing import OptTensor + +class SPINModel(nn.Module): + + def __init__(self, input_size: int, + hidden_size: int, + n_nodes: int, + u_size: Optional[int] = None, + output_size: Optional[int] = None, + temporal_self_attention: bool = True, + reweight: Optional[str] = 'softmax', + n_layers: int = 4, + eta: int = 3, + message_layers: int = 1): + super(SPINModel, self).__init__() + + u_size = u_size or input_size + output_size = output_size or input_size + self.n_nodes = n_nodes + self.n_layers = n_layers + self.eta = eta + self.temporal_self_attention = temporal_self_attention + + self.u_enc = PositionalEncoder(in_channels=u_size, + out_channels=hidden_size, + n_layers=2, + n_nodes=n_nodes) + + self.h_enc = MLP(input_size, hidden_size, n_layers=2) + self.h_norm = LayerNorm(hidden_size) + + self.valid_emb = StaticGraphEmbedding(n_nodes, hidden_size) + self.mask_emb = StaticGraphEmbedding(n_nodes, hidden_size) + + self.x_skip = nn.ModuleList() + self.encoder, self.readout = nn.ModuleList(), nn.ModuleList() + for l in range(n_layers): + x_skip = nn.Linear(input_size, hidden_size) + encoder = TemporalGraphAdditiveAttention( + input_size=hidden_size, + output_size=hidden_size, + msg_size=hidden_size, + msg_layers=message_layers, + temporal_self_attention=temporal_self_attention, + reweight=reweight, + mask_temporal=True, + mask_spatial=l < eta, + norm=True, + root_weight=True, + dropout=0.0 + ) + readout = MLP(hidden_size, hidden_size, output_size, + n_layers=2) + self.x_skip.append(x_skip) + self.encoder.append(encoder) + self.readout.append(readout) + + def forward(self, x: Tensor, u: Tensor, mask: Tensor, + edge_index: Tensor, edge_weight: OptTensor = None, + node_index: OptTensor = None, target_nodes: OptTensor = None): + if target_nodes is None: + target_nodes = slice(None) + + # Whiten missing values + x = x * mask + + # POSITIONAL ENCODING ################################################# + # Obtain spatio-temporal positional encoding for every node-step pair # + # in both observed and target sets. Encoding are obtained by jointly # + # processing node and time positional encoding. # + + # Build (node, timestamp) encoding + q = self.u_enc(u, node_index=node_index) + # Condition value on key + h = self.h_enc(x) + q + + # ENCODER ############################################################# + # Obtain representations h^i_t for every (i, t) node-step pair by # + # only taking into account valid data in representation set. # + + # Replace H in missing entries with queries Q + h = torch.where(mask.bool(), h, q) + # Normalize features + h = self.h_norm(h) + + imputations = [] + + for l in range(self.n_layers): + if l == self.eta: + # Condition H on two different embeddings to distinguish + # valid values from masked ones + valid = self.valid_emb(token_index=node_index) + masked = self.mask_emb(token_index=node_index) + h = torch.where(mask.bool(), h + valid, h + masked) + # Masked Temporal GAT for encoding representation + h = h + self.x_skip[l](x) * mask # skip connection for valid x + #print(f'l={l}', 'h:', tuple(h.shape), 'x:', tuple(x.shape), 'mask:', tuple(mask.shape), 'ei:', edge_index) + h = self.encoder[l](h, edge_index, mask=mask) + # Read from H to get imputations + target_readout = self.readout[l](h[..., target_nodes, :]) + imputations.append(target_readout) + + # Get final layer imputations + x_hat = imputations.pop(-1) + + return x_hat, imputations + + +from inc.diffus import * +from inc.test import * + +import argparse + + +def get_args(): + parser = argparse.ArgumentParser() + # align with existing CLI style (e.g., hermes.py) + parser.add_argument('--dataset', type=str, default=None, + help='dataset name (default: run the built-in list used in the original spin.py)') + parser.add_argument('--seed', type=int, default=123456789, + help='random seed') + parser.add_argument('--data_dir', type=str, default='input', + help='dataset folder') + parser.add_argument('--output', type=str, default='output/spin.pt', + help='output file name') + parser.add_argument('--device', type=torch.device, + default=torch.device('cuda' if torch.cuda.is_available() else 'cpu'), + help='torch device') + + # diffusion parameter estimation (same defaults as the original Dict args) + parser.add_argument('--b_pI0', type=float, default=1e-3, + help='initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type=float, default=1e-3, + help='initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type=int, default=500, + help='optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type=float, default=3e-3, + help='learning rate in diffusion parameter estimation') + + # SPINModel hyperparameters (same defaults as the original Dict args) + parser.add_argument('--u_size', type=int, default=1, + help='size of exogenous input features u') + parser.add_argument('--hidden_size', type=int, default=32, + help='hidden size of SPIN') + parser.add_argument('--temporal_self_attention', dest='temporal_self_attention', + action='store_true', + help='enable temporal self-attention (default: enabled)') + parser.add_argument('--no_temporal_self_attention', dest='temporal_self_attention', + action='store_false', + help='disable temporal self-attention') + parser.set_defaults(temporal_self_attention=True) + parser.add_argument('--reweight', type=str, default='softmax', choices=['softmax', 'l1', 'none'], + help="edge reweighting in attention: 'softmax', 'l1', or 'none'") + parser.add_argument('--n_layers', type=int, default=4, + help='number of layers in SPIN') + parser.add_argument('--eta', type=int, default=3, + help='temporal window size eta in SPIN') + parser.add_argument('--message_layers', type=int, default=1, + help='number of message passing layers in SPIN') + + # Adam + parser.add_argument('--lr', type=float, default=8e-4, + help='learning rate') + parser.add_argument('--l2_reg', type=float, default=0.0, + help='weight decay') + + # training + parser.add_argument('--epochs', type=int, default=300, + help='training epochs') + parser.add_argument('--batch_size', type=int, default=1, + help='number of simulated histories per epoch') + + # evaluation repeats (kept for consistency with other scripts) + parser.add_argument('--rep', type=int, default=1, + help='number of repetitions in evaluation') + + args = parser.parse_args() + + # keep compatibility with the original SPIN code that expects reweight in {'softmax','l1',None} + if getattr(args, 'reweight', None) == 'none': + args.reweight = None + + return args + + +def spin_prep(y, edge_index, args): + n_samples, n_nodes, T = y.size() + T -= 1 + x = y.float().transpose(1, 2).unsqueeze(dim=3) # (samples, T + 1, nodes, 1) + u = torch.ones(n_samples, T + 1, args.u_size, dtype=torch.float32, device=args.device) # (samples, T + 1, u_size) + mask = ( + torch.tensor([[0]] * T + [[1]], device=args.device) + .expand(n_samples, n_nodes, -1, -1) + .transpose(1, 2) + ) # (samples, T + 1, nodes, 1) + ei = edge_index # keep the original behavior + return x, u, mask, ei + + +def spin_run(data, args): + bpar = b_estim(data, args) # diffusion parameter estimation + T = data.T.item() + n_nodes = data.num_nodes + n_out = data.y[:, -1].max().item() + 1 + + # train + model = SPINModel( + input_size=1, + u_size=args.u_size, + n_nodes=n_nodes, + hidden_size=args.hidden_size, + output_size=1, + temporal_self_attention=args.temporal_self_attention, + reweight=args.reweight, + n_layers=args.n_layers, + eta=args.eta, + message_layers=args.message_layers, + ).to(args.device) + + I0 = (data.y[:, 0] == SIR_STATES.I).long().sum().item() + opt = torch.optim.Adam(model.parameters(), lr=args.lr, weight_decay=args.l2_reg) + model.train() + + pbar = trange(1, args.epochs + 1) + for epoch in pbar: + opt.zero_grad() + Y_true = diffus_gen( + T=data.T.item(), + n_nodes=data.num_nodes, + edge_index=data.edge_index, + I0=I0, + n_samples=args.batch_size, + pI=bpar.pI, + pR=bpar.pR, + ) # (T + 1, nodes, samples) + x, u, mask, ei = spin_prep(Y_true.transpose(0, 2), data.edge_index, args) # (samples, T + 1, nodes, *) + z = model(x=x, u=u, mask=mask, edge_index=data.edge_index)[0] # (samples, T + 1, nodes, 1) + loss = (z[:, :-1].flatten() - Y_true[:-1].transpose(1, 2).transpose(0, 1).flatten()).abs().mean() + pbar.set_description(f'[epoch={epoch}] loss={loss.item():.4f}') + loss.backward() + opt.step() + + # infer + with torch.no_grad(): + model.eval() + x, u, mask, ei = spin_prep(data.y.unsqueeze(dim=0).clone(), data.edge_index, args) + z = model(x=x, u=u, mask=mask, edge_index=ei)[0] # (1, T + 1, nodes, 1) + y_pred = data.y.clone() + y_pred[:, :-1] = z[0, :-1, :, 0].clamp(0, n_out - 1).T.round().long() # (nodes, T) + return y_pred.clone() + + +def main(): + args = get_args() + + # Preserve the original default behavior (run a built-in list) when --dataset is not provided. + datasets = ( + ['heb-sir', 'ba-si', 'er-si', 'farmers-si', 'ba-sir', 'er-sir', 'covid-sir'] + if args.dataset is None + else [args.dataset] + ) + + model_fn = lambda data: spin_run(data, args) + tester = Tester(args.data_dir, args.device, model_fn) + tester.test(datasets, seed=args.seed, rep=args.rep) + tester.save(args.output) + + +if __name__ == '__main__': + main() diff --git a/spin_ms.py b/spin_ms.py new file mode 100644 index 0000000..f484b07 --- /dev/null +++ b/spin_ms.py @@ -0,0 +1,709 @@ +# ! pip install --no-index torch-scatter==2.0.7 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +# ! pip install --no-index torch-sparse==0.6.9 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +# ! pip install --no-index torch-cluster==1.5.9 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +#! pip install --no-index torch-spline-conv==1.2.1 -f https://pytorch-geometric.com/whl/torch-1.7.0+cu110.html +# ! pip install torch-geometric==2.0.4 +# ! pip install ndlib==5.1.1 + + +# ! pip install einops==0.6.0 +# ! pip install test_tube==0.7.5 + + +try: + from tsl.nn.base import StaticGraphEmbedding +except Exception: + import torch + from torch import nn + + class StaticGraphEmbedding(nn.Module): + def __init__(self, n_nodes, out_channels): + super().__init__() + self.emb = nn.Embedding(n_nodes, out_channels) + + def forward(self, token_index=None): + if token_index is None: + token_index = torch.arange(self.emb.num_embeddings, device=self.emb.weight.device) + return self.emb(token_index) +from tsl.nn.layers import PositionalEncoding +from tsl.nn.layers.norm import LayerNorm +from tsl.nn.blocks.encoders import MLP +from tsl.nn.functional import sparse_softmax +from tsl.engines import Imputer, Predictor +from tsl.ops.connectivity import weighted_degree +#from tsl.data import Batch, SpatioTemporalDataModule, ImputationDataset + +#SPINModel +'''https://github.com/Graph-Machine-Learning-Group/spin/blob/main/spin/layers/postional_encoding.py''' +from typing import Optional + +from torch import nn + +class PositionalEncoder(nn.Module): + + def __init__(self, in_channels, out_channels, + n_layers: int = 1, + n_nodes: Optional[int] = None): + super(PositionalEncoder, self).__init__() + self.lin = nn.Linear(in_channels, out_channels) + self.activation = nn.LeakyReLU() + self.mlp = MLP(out_channels, out_channels, out_channels, + n_layers=n_layers, activation='relu') + self.positional = PositionalEncoding(out_channels) + if n_nodes is not None: + self.node_emb = StaticGraphEmbedding(n_nodes, out_channels) + else: + self.register_parameter('node_emb', None) + + def forward(self, x, node_emb=None, node_index=None): + if node_emb is None: + node_emb = self.node_emb(token_index=node_index) + # x: [b s c], node_emb: [n c] -> [b s n c] + x = self.lin(x) + x = self.activation(x.unsqueeze(-2) + node_emb) + #print('u:', tuple(x.shape), 'node_emb:', tuple(node_emb.shape))##### + out = self.mlp(x) + out = self.positional(out) + return out + +'''https://github.com/Graph-Machine-Learning-Group/spin/blob/main/spin/layers/additive_attention.py''' +from typing import Optional, Tuple, Union + +import torch +from torch import Tensor +from torch import nn +from torch.nn import LayerNorm, functional as F +from torch_geometric.nn.conv import MessagePassing +from torch_geometric.nn.dense.linear import Linear +from torch_geometric.typing import Adj, OptTensor, PairTensor +from torch_scatter import scatter +from torch_scatter.utils import broadcast + + +class AdditiveAttention(MessagePassing): + def __init__(self, input_size: Union[int, Tuple[int, int]], + output_size: int, + msg_size: Optional[int] = None, + msg_layers: int = 1, + root_weight: bool = True, + reweight: Optional[str] = None, + norm: bool = True, + dropout: float = 0.0, + dim: int = -2, + **kwargs): + kwargs.setdefault('aggr', 'add') + super().__init__(node_dim=dim, **kwargs) + + self.output_size = output_size + if isinstance(input_size, int): + self.src_size = self.tgt_size = input_size + else: + self.src_size, self.tgt_size = input_size + + self.msg_size = msg_size or self.output_size + self.msg_layers = msg_layers + + assert reweight in ['softmax', 'l1', None] + self.reweight = reweight + + self.root_weight = root_weight + self.dropout = dropout + + # key bias is discarded in softmax + self.lin_src = Linear(self.src_size, self.output_size, + weight_initializer='glorot', + bias_initializer='zeros') + self.lin_tgt = Linear(self.tgt_size, self.output_size, + weight_initializer='glorot', bias=False) + + if self.root_weight: + self.lin_skip = Linear(self.tgt_size, self.output_size, + bias=False) + else: + self.register_parameter('lin_skip', None) + + self.msg_nn = nn.Sequential( + nn.PReLU(init=0.2), + MLP(self.output_size, self.msg_size, self.output_size, + n_layers=self.msg_layers, dropout=self.dropout, + activation='prelu') + ) + + if self.reweight == 'softmax': + self.msg_gate = nn.Linear(self.output_size, 1, bias=False) + else: + self.msg_gate = nn.Sequential(nn.Linear(self.output_size, 1), + nn.Sigmoid()) + + if norm: + self.norm = LayerNorm(self.output_size) + else: + self.register_parameter('norm', None) + + self.reset_parameters() + + def reset_parameters(self): + self.lin_src.reset_parameters() + self.lin_tgt.reset_parameters() + if self.lin_skip is not None: + self.lin_skip.reset_parameters() + + def forward(self, x: PairTensor, edge_index: Adj, mask: OptTensor = None): + # if query/key not provided, defaults to x (e.g., for self-attention) + if isinstance(x, Tensor): + x_src = x_tgt = x + else: + x_src, x_tgt = x + x_tgt = x_tgt if x_tgt is not None else x_src + + N_src, N_tgt = x_src.size(self.node_dim), x_tgt.size(self.node_dim) + + msg_src = self.lin_src(x_src) + msg_tgt = self.lin_tgt(x_tgt) + + msg = (msg_src, msg_tgt) + + # propagate_type: (msg: PairTensor, mask: OptTensor) + out = self.propagate(edge_index, msg=msg, mask=mask, + size=(N_src, N_tgt)) + + # skip connection + if self.root_weight: + out = out + self.lin_skip(x_tgt) + + if self.norm is not None: + out = self.norm(out) + + return out + + def normalize_weights(self, weights, index, num_nodes, mask=None): + # mask weights + if mask is not None: + fill_value = float("-inf") if self.reweight == 'softmax' else 0. + weights = weights.masked_fill(torch.logical_not(mask), fill_value) + # eventually reweight + if self.reweight == 'l1': + expanded_index = broadcast(index, weights, self.node_dim) + weights_sum = scatter(weights, expanded_index, self.node_dim, + dim_size=num_nodes, reduce='sum') + weights_sum = weights_sum.index_select(self.node_dim, index) + weights = weights / (weights_sum + 1e-5) + elif self.reweight == 'softmax': + weights = sparse_softmax(weights, index, num_nodes=num_nodes, + dim=self.node_dim) + return weights + + def message(self, msg_j: Tensor, msg_i: Tensor, index, size_i, + mask_j: OptTensor = None) -> Tensor: + msg = self.msg_nn(msg_j + msg_i) + gate = self.msg_gate(msg) + alpha = self.normalize_weights(gate, index, size_i, mask_j) + alpha = F.dropout(alpha, p=self.dropout, training=self.training) + out = alpha * msg + return out + + def __repr__(self) -> str: + return (f'{self.__class__.__name__}({self.output_size}, ' + f'dim={self.node_dim}, ' + f'root_weight={self.root_weight})') + + +class TemporalAdditiveAttention(AdditiveAttention): + def __init__(self, input_size: Union[int, Tuple[int, int]], + output_size: int, + msg_size: Optional[int] = None, + msg_layers: int = 1, + root_weight: bool = True, + reweight: Optional[str] = None, + norm: bool = True, + dropout: float = 0.0, + **kwargs): + kwargs.setdefault('dim', 1) + super().__init__(input_size=input_size, + output_size=output_size, + msg_size=msg_size, + msg_layers=msg_layers, + root_weight=root_weight, + reweight=reweight, + dropout=dropout, + norm=norm, + **kwargs) + + def forward(self, x: PairTensor, mask: OptTensor = None, + temporal_mask: OptTensor = None, + causal_lag: Optional[int] = None): + # x: [b s * c] query: [b l * c] key: [b s * c] + # mask: [b s * c] temporal_mask: [l s] + if isinstance(x, Tensor): + x_src = x_tgt = x + else: + x_src, x_tgt = x + x_tgt = x_tgt if x_tgt is not None else x_src + + l, s = x_tgt.size(self.node_dim), x_src.size(self.node_dim) + i = torch.arange(l, dtype=torch.long, device=x_src.device) + j = torch.arange(s, dtype=torch.long, device=x_src.device) + + # compute temporal index, from j to i + if temporal_mask is None and isinstance(causal_lag, int): + temporal_mask = tuple(torch.tril_indices(l, l, offset=-causal_lag, + device=x_src.device)) + if temporal_mask is not None: + assert temporal_mask.size() == (l, s) + i, j = torch.meshgrid(i, j) + edge_index = torch.stack((j[temporal_mask], i[temporal_mask])) + else: + edge_index = torch.cartesian_prod(j, i).T + + return super(TemporalAdditiveAttention, self).forward(x, edge_index, + mask=mask) + +'''https://github.com/Graph-Machine-Learning-Group/spin/blob/main/spin/layers/temporal_graph_additive_attention.py''' +from typing import Optional, Tuple, Union + +import torch +from torch import Tensor +from torch_geometric.nn.conv import MessagePassing +from torch_geometric.nn.dense.linear import Linear +from torch_geometric.typing import Adj, OptTensor, OptPairTensor + +class TemporalGraphAdditiveAttention(MessagePassing): + def __init__(self, input_size: Union[int, Tuple[int, int]], + output_size: int, + msg_size: Optional[int] = None, + msg_layers: int = 1, + root_weight: bool = True, + reweight: Optional[str] = None, + temporal_self_attention: bool = True, + mask_temporal: bool = True, + mask_spatial: bool = True, + norm: bool = True, + dropout: float = 0., + **kwargs): + kwargs.setdefault('aggr', 'add') + super(TemporalGraphAdditiveAttention, self).__init__(node_dim=-2, + **kwargs) + + # store dimensions + if isinstance(input_size, int): + self.src_size = self.tgt_size = input_size + else: + self.src_size, self.tgt_size = input_size + self.output_size = output_size + self.msg_size = msg_size or self.output_size + + self.mask_temporal = mask_temporal + self.mask_spatial = mask_spatial + + self.root_weight = root_weight + self.dropout = dropout + + if temporal_self_attention: + self.self_attention = TemporalAdditiveAttention( + input_size=input_size, + output_size=output_size, + msg_size=msg_size, + msg_layers=msg_layers, + reweight=reweight, + dropout=dropout, + root_weight=False, + norm=False + ) + else: + self.register_parameter('self_attention', None) + + self.cross_attention = TemporalAdditiveAttention(input_size=input_size, + output_size=output_size, + msg_size=msg_size, + msg_layers=msg_layers, + reweight=reweight, + dropout=dropout, + root_weight=False, + norm=False) + + if self.root_weight: + self.lin_skip = Linear(self.tgt_size, self.output_size, + bias_initializer='zeros') + else: + self.register_parameter('lin_skip', None) + + if norm: + self.norm = LayerNorm(output_size) + else: + self.register_parameter('norm', None) + + self.reset_parameters() + + def reset_parameters(self): + self.cross_attention.reset_parameters() + if self.self_attention is not None: + self.self_attention.reset_parameters() + if self.lin_skip is not None: + self.lin_skip.reset_parameters() + if self.norm is not None: + self.norm.reset_parameters() + + def forward(self, x: OptPairTensor, + edge_index: Adj, edge_weight: OptTensor = None, + mask: OptTensor = None): + # inputs: [batch, steps, nodes, channels] + if isinstance(x, Tensor): + x_src = x_tgt = x + else: + x_src, x_tgt = x + x_tgt = x_tgt if x_tgt is not None else x_src + + n_src, n_tgt = x_src.size(-2), x_tgt.size(-2) + + # propagate query, key and value + #print('src:', x_src.shape, 'tgt:', x_tgt.shape, 'ei:', edge_index.shape, 'mask:', mask.shape, f'mask_spatial={self.mask_spatial}') + out = self.propagate(x=(x_src, x_tgt), + edge_index=edge_index, edge_weight=edge_weight, + mask=mask if self.mask_spatial else None, + size=(n_src, n_tgt)) + + if self.self_attention is not None: + s, l = x_src.size(1), x_tgt.size(1) + if s == l: + attn_mask = ~torch.eye(l, l, dtype=torch.bool, + device=x_tgt.device) + else: + attn_mask = None + temp = self.self_attention(x=(x_src, x_tgt), + mask=mask if self.mask_temporal else None, + temporal_mask=attn_mask) + out = out + temp + + # skip connection + if self.root_weight: + out = out + self.lin_skip(x_tgt) + + if self.norm is not None: + out = self.norm(out) + + return out + + def message(self, x_i: Tensor, x_j: Tensor, + edge_weight: OptTensor, mask_j: OptTensor) -> Tensor: + # [batch, steps, edges, channels] + + out = self.cross_attention((x_j, x_i), mask=mask_j) + #print('out:', out.shape) + + if edge_weight is not None: + out = out * edge_weight.view(-1, 1) + return out + +'''https://github.com/Graph-Machine-Learning-Group/spin/blob/main/spin/models/spin.py''' +from typing import Optional + +import torch +from torch import nn, Tensor +from torch.nn import LayerNorm +from torch_geometric.typing import OptTensor + +class SPINModel(nn.Module): + + def __init__(self, input_size: int, + hidden_size: int, + n_nodes: int, + u_size: Optional[int] = None, + output_size: Optional[int] = None, + temporal_self_attention: bool = True, + reweight: Optional[str] = 'softmax', + n_layers: int = 4, + eta: int = 3, + message_layers: int = 1): + super(SPINModel, self).__init__() + + u_size = u_size or input_size + output_size = output_size or input_size + self.n_nodes = n_nodes + self.n_layers = n_layers + self.eta = eta + self.temporal_self_attention = temporal_self_attention + + self.u_enc = PositionalEncoder(in_channels=u_size, + out_channels=hidden_size, + n_layers=2, + n_nodes=n_nodes) + + self.h_enc = MLP(input_size, hidden_size, n_layers=2) + self.h_norm = LayerNorm(hidden_size) + + self.valid_emb = StaticGraphEmbedding(n_nodes, hidden_size) + self.mask_emb = StaticGraphEmbedding(n_nodes, hidden_size) + + self.x_skip = nn.ModuleList() + self.encoder, self.readout = nn.ModuleList(), nn.ModuleList() + for l in range(n_layers): + x_skip = nn.Linear(input_size, hidden_size) + encoder = TemporalGraphAdditiveAttention( + input_size=hidden_size, + output_size=hidden_size, + msg_size=hidden_size, + msg_layers=message_layers, + temporal_self_attention=temporal_self_attention, + reweight=reweight, + mask_temporal=True, + mask_spatial=l < eta, + norm=True, + root_weight=True, + dropout=0.0 + ) + readout = MLP(hidden_size, hidden_size, output_size, + n_layers=2) + self.x_skip.append(x_skip) + self.encoder.append(encoder) + self.readout.append(readout) + + def forward(self, x: Tensor, u: Tensor, mask: Tensor, + edge_index: Tensor, edge_weight: OptTensor = None, + node_index: OptTensor = None, target_nodes: OptTensor = None): + if target_nodes is None: + target_nodes = slice(None) + + # Whiten missing values + x = x * mask + + # POSITIONAL ENCODING ################################################# + # Obtain spatio-temporal positional encoding for every node-step pair # + # in both observed and target sets. Encoding are obtained by jointly # + # processing node and time positional encoding. # + + # Build (node, timestamp) encoding + q = self.u_enc(u, node_index=node_index) + # Condition value on key + h = self.h_enc(x) + q + + # ENCODER ############################################################# + # Obtain representations h^i_t for every (i, t) node-step pair by # + # only taking into account valid data in representation set. # + + # Replace H in missing entries with queries Q + h = torch.where(mask.bool(), h, q) + # Normalize features + h = self.h_norm(h) + + imputations = [] + + for l in range(self.n_layers): + if l == self.eta: + # Condition H on two different embeddings to distinguish + # valid values from masked ones + valid = self.valid_emb(token_index=node_index) + masked = self.mask_emb(token_index=node_index) + h = torch.where(mask.bool(), h + valid, h + masked) + # Masked Temporal GAT for encoding representation + h = h + self.x_skip[l](x) * mask # skip connection for valid x + #print(f'l={l}', 'h:', tuple(h.shape), 'x:', tuple(x.shape), 'mask:', tuple(mask.shape), 'ei:', edge_index) + h = self.encoder[l](h, edge_index, mask=mask) + # Read from H to get imputations + target_readout = self.readout[l](h[..., target_nodes, :]) + imputations.append(target_readout) + + # Get final layer imputations + x_hat = imputations.pop(-1) + + return x_hat, imputations + + +import argparse +import torch + +from inc.diffus import * + +try: + from inc.test_ms import Tester # type: ignore +except Exception: + from inc.test import Tester # type: ignore + + +def get_args(): + parser = argparse.ArgumentParser() + + # -------------------- common experiment args -------------------- + parser.add_argument('--dataset', type=str, required=True, help='dataset name') + parser.add_argument('--seed', type=int, default=123456789, help='random seed') + parser.add_argument('--data_dir', type=str, default='input', help='dataset folder') + parser.add_argument('--output', type=str, default='output/spin.pt', help='output file name') + parser.add_argument( + '--device', type=str, + default='cuda' if torch.cuda.is_available() else 'cpu', + help='torch device, e.g., "cuda", "cuda:0", or "cpu"' + ) + parser.add_argument( + '--obs_time', '--snapshot', dest='obs_time', + type=str, default='', + help='extra observed snapshot times, comma-separated, e.g., "5,7,9". ' + 'Final time T will always be added automatically.' + ) + + # -------------------- diffusion parameter estimation -------------------- + parser.add_argument('--b_pI0', type=float, default=1e-3, + help='initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type=float, default=1e-3, + help='initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type=int, default=500, + help='optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type=float, default=3e-3, + help='learning rate in diffusion parameter estimation') + + # -------------------- SPIN model hyperparams -------------------- + parser.add_argument('--u_size', type=int, default=1, help='u feature size (kept as 1)') + parser.add_argument('--hidden_size', type=int, default=32, help='hidden size') + parser.add_argument('--reweight', type=str, default='softmax', + choices=['softmax', 'l1', 'none'], + help='attention reweighting: softmax | l1 | none') + parser.add_argument('--n_layers', type=int, default=4, help='number of SPIN layers') + parser.add_argument('--eta', type=int, default=3, help='layers before enabling spatial mask-off') + parser.add_argument('--message_layers', type=int, default=1, help='message MLP depth inside attention') + + parser.add_argument( + '--no_temporal_self_attention', action='store_false', + dest='temporal_self_attention', + help='disable temporal self-attention (default: enabled)' + ) + parser.set_defaults(temporal_self_attention=True) + + # -------------------- training hyperparams -------------------- + parser.add_argument('--lr', type=float, default=8e-4, help='Adam lr') + parser.add_argument('--l2_reg', type=float, default=0.0, help='Adam weight_decay') + parser.add_argument('--epochs', type=int, default=300, help='training epochs') + parser.add_argument('--batch_size', type=int, default=1, help='synthetic batch size per epoch') + + args = parser.parse_args() + args.device = torch.device(args.device) + + if args.reweight == 'none': + args.reweight = None + + return args + + +def parse_obs_time(obs_time_str: str, T: int): + """Parse comma-separated observed snapshot times and always include final time T.""" + times = [] + if obs_time_str: + for part in str(obs_time_str).split(','): + part = part.strip() + if part == '': + continue + times.append(int(part)) + # keep within [0, T] + times = [t for t in times if 0 <= t <= T] + if T not in times: + times.append(T) + return sorted(set(times)) + + +def build_obs_mask(n_samples: int, n_nodes: int, T: int, obs_time, device): + """Return float mask with shape (samples, T+1, nodes, 1), 1=observed, 0=missing.""" + mask = torch.zeros((n_samples, T + 1, n_nodes, 1), dtype=torch.float32, device=device) + if len(obs_time) > 0: + mask[:, obs_time, :, :] = 1.0 + return mask + + +def spin_prep(y, edge_index, obs_time, args): + """ + y: (samples, nodes, T+1) + return: + x: (samples, T+1, nodes, 1) + u: (samples, T+1, u_size) + mask: (samples, T+1, nodes, 1) + ei: edge_index + """ + n_samples, n_nodes, Tp1 = y.size() + T = Tp1 - 1 + + x = y.float().transpose(1, 2).unsqueeze(dim=3) # (samples, T+1, nodes, 1) + u = torch.ones(n_samples, T + 1, args.u_size, dtype=torch.float32, device=args.device) + mask = build_obs_mask(n_samples, n_nodes, T, obs_time, device=args.device) + + # NOTE: SPIN implementation here uses a single static edge_index shared across all steps. + ei = edge_index + return x, u, mask, ei + + +def spin_run(data, args): + """Train SPIN on synthetic histories and infer missing diffusion history.""" + T = int(data.T.item()) + obs_time = parse_obs_time(args.obs_time, T) + data.obs_ts = obs_time + bpar = b_estim(data, args, obs_time=obs_time) + + n_nodes = int(data.num_nodes) + n_out = int(data.y[:, -1].max().item() + 1) + + # -------------------- train -------------------- + model = SPINModel( + input_size=1, + u_size=args.u_size, + n_nodes=n_nodes, + hidden_size=args.hidden_size, + output_size=1, + temporal_self_attention=args.temporal_self_attention, + reweight=args.reweight, + n_layers=args.n_layers, + eta=args.eta, + message_layers=args.message_layers, + ).to(args.device) + + I0 = int((data.y[:, 0] == SIR_STATES.I).long().sum().item()) + opt = torch.optim.Adam(model.parameters(), lr=args.lr, weight_decay=args.l2_reg) + + model.train() + pbar = trange(1, args.epochs + 1) + for epoch in pbar: + opt.zero_grad() + + # synthetic history + Y_true = diffus_gen( + T=T, + n_nodes=n_nodes, + edge_index=data.edge_index, + I0=I0, + n_samples=args.batch_size, + pI=bpar.pI, + pR=bpar.pR, + ) # (T+1, nodes, samples) + + x, u, mask, ei = spin_prep(Y_true.transpose(0, 2), data.edge_index, obs_time, args) + z = model(x=x, u=u, mask=mask, edge_index=ei)[0] # (samples, T+1, nodes, 1) + y_true = Y_true.permute(2, 0, 1).unsqueeze(-1) # (samples, T+1, nodes, 1) + unobs = ~mask.bool() + loss = (z - y_true).abs() + loss = loss[unobs].mean() if unobs.any() else loss.mean() + + pbar.set_description(f'[epoch={epoch}] loss={loss.item():.4f}') + loss.backward() + opt.step() + + # -------------------- infer -------------------- + with torch.no_grad(): + model.eval() + + x, u, mask, ei = spin_prep(data.y.unsqueeze(0).clone(), data.edge_index, obs_time, args) + z = model(x=x, u=u, mask=mask, edge_index=ei)[0] # (1, T+1, nodes, 1) + + y_pred = z[0, :, :, 0].clamp(0, n_out - 1).transpose(0, 1).round().long() # (nodes, T+1) + for t in obs_time: + y_pred[:, t] = data.y[:, t] + + return y_pred.detach().clone() + + +def main(): + args = get_args() + + def _run(data): + return spin_run(data, args) + + tester = Tester(args.data_dir, args.device, _run) + tester.test([args.dataset], seed=args.seed, rep=1) + tester.save(args.output) + + +if __name__ == '__main__': + main() From 00cb433b05ff40d32fe9dbee842543b44dc746f2 Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Thu, 29 Jan 2026 17:58:53 -0600 Subject: [PATCH 07/19] fix the tester problem --- brits_ms.py | 7 +------ ditto.py | 11 ++++++++++- grin_ms.py | 4 +--- spin_ms.py | 5 +---- 4 files changed, 13 insertions(+), 14 deletions(-) diff --git a/brits_ms.py b/brits_ms.py index 65bb174..1b48825 100644 --- a/brits_ms.py +++ b/brits_ms.py @@ -13,12 +13,7 @@ from torch.autograd import Variable from torch.nn.parameter import Parameter -try: - # Multi-snapshot tester (recommended). - from inc.test_ms import Tester -except Exception: - # Fall back to single-snapshot tester if needed. - from inc.test import Tester +from inc.test import Tester from tqdm import trange diff --git a/ditto.py b/ditto.py index b4e482c..fb86808 100644 --- a/ditto.py +++ b/ditto.py @@ -121,7 +121,16 @@ def lik(self, Y): # Y: (T+1, nodes, samples) rem.flatten()[vids] = torch.where(mski.repeat_interleave(repeats = degi), torch.where(trsi.repeat_interleave(repeats = degi), rems - 1, self.n_inf), rems) # (T * sum neighbs) mskI[uidi] &= opti # (T * samples) # likR + likI - lik = (torch.where(mskR, torch.where(trsR, lR1, lR0), self.zero).view(-1, n_samples) + torch.where(mskI.view(-1, n_samples), torch.where(trsI.view(-1, n_samples), lI1.view(-1, n_samples), lI0.view(-1, n_samples)), self.zero)).sum(dim = 0) # (samples,) + likR = torch.where(mskR, torch.where(trsR, lR1, lR0), self.zero).reshape(-1, n_samples) + + mskI2 = mskI.reshape(-1, n_samples) + trsI2 = trsI.reshape(-1, n_samples) + lI1_2 = lI1.reshape(-1, n_samples) + lI0_2 = lI0.reshape(-1, n_samples) + + likI = torch.where(mskI2, torch.where(trsI2, lI1_2, lI0_2), self.zero) + + lik = (likR + likI).sum(dim=0) return lik, zI0, zR0, zI, zR # (samples,) @torch.no_grad() def clamp_grad(self, z0, grad): diff --git a/grin_ms.py b/grin_ms.py index efb624a..afc7415 100644 --- a/grin_ms.py +++ b/grin_ms.py @@ -12,9 +12,7 @@ # Project utilities (same style as other runners) from inc.diffus import SIR_STATES, b_estim, diffus_gen, seed_all -# Multi-snapshot aware tester (only evaluates on unobserved positions) -# If your repo uses inc.test as the ms tester, you can swap this import accordingly. -from inc.test_ms import Tester +from inc.test import Tester def build_parser() -> argparse.ArgumentParser: diff --git a/spin_ms.py b/spin_ms.py index f484b07..4df7c37 100644 --- a/spin_ms.py +++ b/spin_ms.py @@ -513,10 +513,7 @@ def forward(self, x: Tensor, u: Tensor, mask: Tensor, from inc.diffus import * -try: - from inc.test_ms import Tester # type: ignore -except Exception: - from inc.test import Tester # type: ignore +from inc.test import Tester # type: ignore def get_args(): From 3f84bb4f610bbbc2bb9be08981d14647ab3ac98d Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Thu, 29 Jan 2026 18:54:52 -0600 Subject: [PATCH 08/19] print accept rate --- hermes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/hermes.py b/hermes.py index 44b6921..febe412 100644 --- a/hermes.py +++ b/hermes.py @@ -311,7 +311,7 @@ def t_mcmc(data, bpar, q_net, args, obs_time, keepdim=True): # Hastings acceptance a = torch.rand(args.t_samples, device=args.device) <= torch.exp(lpY + lqX - lpX - lqY) - + pbar.set_description(f"[step={step}] acc={a.float().mean().item():.3f}") X = torch.where(a, Y, X) lqX = torch.where(a, lqY, lqX) lpX = torch.where(a, lpY, lpX) From 8b6f96d157d48c743e962cc28dc668444a37100b Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Mon, 2 Feb 2026 23:15:15 -0600 Subject: [PATCH 09/19] update 2 experiments --- experiments/exp1_scalability.py | 511 ++++++++++++++++++ experiments/exp2_timespan.py | 317 +++++++++++ hermes.py | 13 +- .../exp1_scalability/ba-si-n1000-T1-g0-d1.pt | Bin 0 -> 178271 bytes .../vsN/farmers-si_n82_T16.pt | Bin 0 -> 21899 bytes .../exp1_scalability/vsT/farmers-si_n82_T4.pt | Bin 0 -> 14017 bytes .../exp1_scalability/vsT/farmers-si_n82_T6.pt | Bin 0 -> 15297 bytes .../exp1_scalability/vsT/farmers-si_n82_T8.pt | Bin 0 -> 16641 bytes input/exp2_timespan/T3/synthetic/ba-sir.pt | Bin 0 -> 178126 bytes 9 files changed, 837 insertions(+), 4 deletions(-) create mode 100644 experiments/exp1_scalability.py create mode 100644 experiments/exp2_timespan.py create mode 100644 input/exp1_scalability/ba-si-n1000-T1-g0-d1.pt create mode 100644 input/exp1_scalability/vsN/farmers-si_n82_T16.pt create mode 100644 input/exp1_scalability/vsT/farmers-si_n82_T4.pt create mode 100644 input/exp1_scalability/vsT/farmers-si_n82_T6.pt create mode 100644 input/exp1_scalability/vsT/farmers-si_n82_T8.pt create mode 100644 input/exp2_timespan/T3/synthetic/ba-sir.pt diff --git a/experiments/exp1_scalability.py b/experiments/exp1_scalability.py new file mode 100644 index 0000000..d39e655 --- /dev/null +++ b/experiments/exp1_scalability.py @@ -0,0 +1,511 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +Experiment 1: Scalability (HERMES) + +Runtime profiling for: + (1) vs T (fixed n, vary T) + (2) vs n (fixed T, vary n) + +Dataset: ba-sir +Observation: only TWO observed frames for each run: + obs_time = [floor(T/2), T] + +Defaults for vsT: + T in {3,4,5,6,7,8,9,10} (skip T=1,2) + +Defaults for vsN: + T_fixed = 10 + scale n by tiling disjoint copies (factors) + +Run from repo root: + python experiments/exp1_scalability.py --method hermes --dataset ba-sir --data_dir input --device cuda + +Optional: + --save_datasets to export generated .pt datasets (for traceability) +""" + +from __future__ import annotations + +import os +import sys +import gc +import csv +import time +import argparse +import platform +from dataclasses import dataclass, asdict +from datetime import datetime +from typing import List, Dict, Any + +import torch +from torch_geometric.data import Data + +# --------------------------------------------------------------------- +# Make repo root importable (so `import inc.*` and `import hermes` works) +# --------------------------------------------------------------------- +ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) +if ROOT not in sys.path: + sys.path.insert(0, ROOT) + +# project imports (re-use existing code) +from inc.data import data_load, data_make_states # type: ignore +import hermes # hermes.py must be import-safe (guarded by __main__) + +# seeding util (fallback if inc.utils not present) +try: + from inc.utils import seed_all # type: ignore +except Exception: # pragma: no cover + import random + import numpy as np + + def seed_all(seed: int) -> None: + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +CSV_FIELDS = [ + "timestamp", + "exp", + "method", + "base_dataset", + "dataset_pt", + "seed", + "device", + "status", + "n_nodes", + "n_edges", + "T", + "obs_time_mid", + "obs_time", + "n_factor", + # HERMES HP + "b_pI0", + "b_pR0", + "b_steps", + "b_lr", + "q_steps", + "q_lr", + "q_hid", + "q_gnn", + "q_mlp", + "q_samples", + "q_zlim", + "p_coef", + "t_samples", + "t_steps", + "t_keep", + # timings + "b_estim_sec", + "q_train_sec", + "t_mcmc_sec", + "total_sec", + # meta + "host", +] + + +def _sync(device: torch.device) -> None: + if device.type == "cuda": + torch.cuda.synchronize(device) + + +def _parse_int_list(s: str) -> List[int]: + if s is None: + return [] + s = s.strip() + if not s: + return [] + parts = s.replace(",", " ").split() + return [int(x) for x in parts] + + +def _ensure_dir(path: str) -> None: + os.makedirs(path, exist_ok=True) + + +def _safe_save_pt(data: Data, path: str) -> None: + _ensure_dir(os.path.dirname(path)) + torch.save(data.cpu(), path) + + +def _variant_timespan(base: Data, T_new: int) -> Data: + """ + Truncate/clip diffusion timespan to T_new and recompute y from (clipped) tI/tR. + """ + assert T_new >= 1, "T_new should be >= 1" + device = base.tI.device + T_tensor = torch.tensor(T_new, dtype=base.T.dtype, device=device) + + # Clip hitting times beyond T_new to 'never' = T_new+1 + cap = torch.full_like(base.tI, T_new + 1) + tI_new = torch.minimum(base.tI, cap) + + tR_new = None + if hasattr(base, "tR") and getattr(base, "tR") is not None: + capR = torch.full_like(base.tR, T_new + 1) + tR_new = torch.minimum(base.tR, capR) + + y_new = data_make_states(T_new, tI_new, tR_new) + + out = Data(edge_index=base.edge_index.clone(), y=y_new, T=T_tensor, tI=tI_new.clone()) + if tR_new is not None: + out.tR = tR_new.clone() + + out.num_nodes = y_new.size(0) + return out + + +def _variant_scale_n(base: Data, factor: int) -> Data: + """ + Tile disjoint copies of a base graph/history to scale n. + This avoids rewriting the synthetic generator and is sufficient for runtime scaling. + """ + assert factor >= 1 + if factor == 1: + out = Data(edge_index=base.edge_index.clone(), y=base.y.clone(), T=base.T.clone(), tI=base.tI.clone()) + if hasattr(base, "tR") and getattr(base, "tR") is not None: + out.tR = base.tR.clone() + out.num_nodes = base.num_nodes + return out + + n0 = int(base.num_nodes) + eidx_list = [base.edge_index + k * n0 for k in range(factor)] + edge_index = torch.cat(eidx_list, dim=1) + + y = base.y.repeat(factor, 1) + tI = base.tI.repeat(factor) + out = Data(edge_index=edge_index, y=y, T=base.T.clone(), tI=tI) + if hasattr(base, "tR") and getattr(base, "tR") is not None: + out.tR = base.tR.repeat(factor) + + out.num_nodes = y.size(0) + return out + + +@dataclass +class HermesHP: + """ + Defaults match your provided ba-sir command: + --b_pI0 0.001 --b_pR0 0.001 --b_steps 500 --b_lr 0.003 + --q_steps 250 --q_lr 0.003 --q_hid 16 --q_gnn 3 --q_mlp 2 --q_samples 10 --q_zlim 16 + --p_coef 1.0 --t_samples 100 --t_steps 100 --t_keep 0.5 + """ + b_pI0: float = 0.001 + b_pR0: float = 0.001 + b_steps: int = 500 + b_lr: float = 0.003 + q_steps: int = 250 + q_lr: float = 0.003 + q_hid: int = 16 + q_gnn: int = 3 + q_mlp: int = 2 + q_samples: int = 10 + q_zlim: int = 16 + p_coef: float = 1.0 + t_samples: int = 100 + t_steps: int = 100 + t_keep: float = 0.5 + + def to_namespace(self, *, device: torch.device, seed: int, obs_time_mid: int) -> argparse.Namespace: + ns = argparse.Namespace(**asdict(self)) + ns.device = device + ns.seed = seed + # for compatibility; we pass obs_time explicitly in calls anyway + ns.obs_time = str(obs_time_mid) + return ns + + +def _run_hermes_once_timed(data: Data, args: argparse.Namespace, obs_time: List[int]) -> Dict[str, float]: + """ + Run HERMES pipeline once and return per-stage runtimes (seconds): + b_estim, q_train, t_mcmc, total + """ + obs_time = sorted(set(int(t) for t in obs_time)) + + # move to device (exclude copy time from timing) + data = data.to(args.device) + + _sync(args.device) + t0 = time.perf_counter() + + # 1) diffusion parameter estimation + b0 = time.perf_counter() + bpar = hermes.b_estim(data, args, obs_time=obs_time) + _sync(args.device) + t_b = time.perf_counter() - b0 + + # 2) proposal training + q0 = time.perf_counter() + q_net = hermes.q_train(data, obs_time, bpar, args) + _sync(args.device) + t_q = time.perf_counter() - q0 + + # 3) MCMC inference + m0 = time.perf_counter() + _ = hermes.t_mcmc(data, bpar, q_net, args, obs_time=obs_time, keepdim=True) + _sync(args.device) + t_m = time.perf_counter() - m0 + + t_total = time.perf_counter() - t0 + + # cleanup (outside timing) + del q_net, bpar + gc.collect() + if args.device.type == "cuda": + torch.cuda.empty_cache() + + return {"b_estim_sec": t_b, "q_train_sec": t_q, "t_mcmc_sec": t_m, "total_sec": t_total} + + +def _append_csv(path: str, row: Dict[str, Any]) -> None: + _ensure_dir(os.path.dirname(path)) + file_exists = os.path.exists(path) + with open(path, "a", newline="", encoding="utf-8") as f: + writer = csv.DictWriter(f, fieldnames=CSV_FIELDS, extrasaction="ignore") + if not file_exists: + writer.writeheader() + fixed_row = {k: row.get(k, "") for k in CSV_FIELDS} + writer.writerow(fixed_row) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--method", type=str, default="hermes", choices=["hermes"]) + parser.add_argument("--dataset", type=str, default="ba-sir") + parser.add_argument("--data_dir", type=str, default="input") + parser.add_argument("--device", type=str, default="cuda") + parser.add_argument("--seed", type=int, default=123456789) + + # Exp controls + parser.add_argument("--run_vsT", action="store_true", help="run scalability vs T") + parser.add_argument("--run_vsN", action="store_true", help="run scalability vs n") + + # ba-sir: T=3..10 (skip 1,2) + parser.add_argument("--T_list", type=str, default="3,4,5,6,7,8,9,10", + help="comma-separated T list for vsT") + parser.add_argument("--T_fixed", type=int, default=10, help="fixed T for vsN") + + # scale-n factors (n = factor * n0). You can override to larger factors, script will mark OOM if happens. + parser.add_argument("--n_factors", type=str, default="1,2,4,8,16", + help="tile factors for vsN (n = factor * n0)") + parser.add_argument("--n_factor_for_vsT", type=int, default=1, + help="optionally scale n before varying T") + + # IO + parser.add_argument("--gen_dir", type=str, default="", + help="where to save generated .pt datasets (default: /exp1_scalability)") + parser.add_argument("--save_datasets", action="store_true", help="save generated datasets as .pt") + parser.add_argument("--out_csv", type=str, default="output/exp1_scalability_runtime_ba_sir.csv") + parser.add_argument("--overwrite_csv", action="store_true") + + # HERMES hyperparameters (defaults match your provided ba-sir command) + parser.add_argument("--b_pI0", type=float, default=0.001) + parser.add_argument("--b_pR0", type=float, default=0.001) + parser.add_argument("--b_steps", type=int, default=500) + parser.add_argument("--b_lr", type=float, default=0.003) + parser.add_argument("--q_steps", type=int, default=250) + parser.add_argument("--q_lr", type=float, default=0.003) + parser.add_argument("--q_hid", type=int, default=16) + parser.add_argument("--q_gnn", type=int, default=3) + parser.add_argument("--q_mlp", type=int, default=2) + parser.add_argument("--q_samples", type=int, default=10) + parser.add_argument("--q_zlim", type=int, default=16) + parser.add_argument("--p_coef", type=float, default=1.0) + parser.add_argument("--t_samples", type=int, default=100) + parser.add_argument("--t_steps", type=int, default=100) + parser.add_argument("--t_keep", type=float, default=0.5) + + args_cli = parser.parse_args() + + if not args_cli.run_vsT and not args_cli.run_vsN: + # default: run both + args_cli.run_vsT = True + args_cli.run_vsN = True + + # device resolve + if args_cli.device.startswith("cuda") and not torch.cuda.is_available(): + print("[warn] CUDA not available, fallback to CPU.") + device = torch.device("cpu") + else: + device = torch.device(args_cli.device) + + # output csv + if args_cli.overwrite_csv and os.path.exists(args_cli.out_csv): + os.remove(args_cli.out_csv) + + # gen dir + gen_dir = args_cli.gen_dir.strip() or os.path.join(args_cli.data_dir, "exp1_scalability") + _ensure_dir(gen_dir) + + # load base dataset on CPU (generation happens on CPU, then move to device for timing) + base = data_load(args_cli.dataset, args_cli.data_dir, torch.device("cpu")) + n0 = int(base.num_nodes) + m0 = int(base.edge_index.size(1)) + T0 = int(base.T.item()) + print(f"[load] {args_cli.dataset}: n={n0}, m={m0}, T={T0}") + + # build HP template + hp = HermesHP( + b_pI0=args_cli.b_pI0, + b_pR0=args_cli.b_pR0, + b_steps=args_cli.b_steps, + b_lr=args_cli.b_lr, + q_steps=args_cli.q_steps, + q_lr=args_cli.q_lr, + q_hid=args_cli.q_hid, + q_gnn=args_cli.q_gnn, + q_mlp=args_cli.q_mlp, + q_samples=args_cli.q_samples, + q_zlim=args_cli.q_zlim, + p_coef=args_cli.p_coef, + t_samples=args_cli.t_samples, + t_steps=args_cli.t_steps, + t_keep=args_cli.t_keep, + ) + + host = platform.node() + + # ------------------------- + # run vsT + # ------------------------- + if args_cli.run_vsT: + T_list = _parse_int_list(args_cli.T_list) + assert len(T_list) > 0, "T_list is empty" + print(f"[exp] vsT: T_list={T_list}, n_factor_for_vsT={args_cli.n_factor_for_vsT}") + + base_scaled = _variant_scale_n(base, args_cli.n_factor_for_vsT) + + for T in T_list: + if T > T0: + print(f"[skip] T={T} > base T0={T0}.") + continue + if T <= 2: + print(f"[skip] T={T} (skip T<=2 for this setting).") + continue + + data_T = _variant_timespan(base_scaled, T) + obs_mid = T // 2 + obs_time = [obs_mid, T] # exactly TWO frames + + pt_path = "" + if args_cli.save_datasets: + pt_path = os.path.join(gen_dir, "vsT", f"{args_cli.dataset}_n{data_T.num_nodes}_T{T}.pt") + _safe_save_pt(data_T, pt_path) + + seed_all(args_cli.seed) + run_args = hp.to_namespace(device=device, seed=args_cli.seed, obs_time_mid=obs_mid) + + status = "ok" + times = {"b_estim_sec": float("nan"), "q_train_sec": float("nan"), + "t_mcmc_sec": float("nan"), "total_sec": float("nan")} + try: + times = _run_hermes_once_timed(data_T, run_args, obs_time) + except RuntimeError as e: + if "out of memory" in str(e).lower(): + status = "oom" + if device.type == "cuda": + torch.cuda.empty_cache() + else: + raise + + row = { + "timestamp": datetime.now().isoformat(timespec="seconds"), + "exp": "vsT", + "method": args_cli.method, + "base_dataset": args_cli.dataset, + "dataset_pt": pt_path, + "seed": args_cli.seed, + "device": str(device), + "status": status, + "n_nodes": int(data_T.num_nodes), + "n_edges": int(data_T.edge_index.size(1)), + "T": int(T), + "obs_time_mid": int(obs_mid), + "obs_time": ",".join(map(str, obs_time)), + "n_factor": int(args_cli.n_factor_for_vsT), + **asdict(hp), + **times, + "host": host, + } + _append_csv(args_cli.out_csv, row) + print(f"[done] vsT T={T} n={row['n_nodes']} total={row['total_sec']:.3f}s status={status}") + + del data_T + gc.collect() + + # ------------------------- + # run vsN + # ------------------------- + if args_cli.run_vsN: + factors = _parse_int_list(args_cli.n_factors) + assert len(factors) > 0, "n_factors is empty" + T_fixed = int(args_cli.T_fixed) + if T_fixed > T0: + print(f"[warn] T_fixed={T_fixed} > base T0={T0}, truncate to T0={T0}.") + T_fixed = T0 + + print(f"[exp] vsN: factors={factors}, T_fixed={T_fixed}") + + base_T = _variant_timespan(base, T_fixed) + obs_mid = T_fixed // 2 + obs_time = [obs_mid, T_fixed] # exactly TWO frames + + for fac in factors: + data_N = _variant_scale_n(base_T, fac) + + pt_path = "" + if args_cli.save_datasets: + pt_path = os.path.join(gen_dir, "vsN", f"{args_cli.dataset}_n{data_N.num_nodes}_T{T_fixed}.pt") + _safe_save_pt(data_N, pt_path) + + seed_all(args_cli.seed) + run_args = hp.to_namespace(device=device, seed=args_cli.seed, obs_time_mid=obs_mid) + + status = "ok" + times = {"b_estim_sec": float("nan"), "q_train_sec": float("nan"), + "t_mcmc_sec": float("nan"), "total_sec": float("nan")} + try: + times = _run_hermes_once_timed(data_N, run_args, obs_time) + except RuntimeError as e: + if "out of memory" in str(e).lower(): + status = "oom" + if device.type == "cuda": + torch.cuda.empty_cache() + else: + raise + + row = { + "timestamp": datetime.now().isoformat(timespec="seconds"), + "exp": "vsN", + "method": args_cli.method, + "base_dataset": args_cli.dataset, + "dataset_pt": pt_path, + "seed": args_cli.seed, + "device": str(device), + "status": status, + "n_nodes": int(data_N.num_nodes), + "n_edges": int(data_N.edge_index.size(1)), + "T": int(T_fixed), + "obs_time_mid": int(obs_mid), + "obs_time": ",".join(map(str, obs_time)), + "n_factor": int(fac), + **asdict(hp), + **times, + "host": host, + } + _append_csv(args_cli.out_csv, row) + print(f"[done] vsN fac={fac} n={row['n_nodes']} total={row['total_sec']:.3f}s status={status}") + + del data_N + gc.collect() + + print(f"[ok] wrote CSV -> {args_cli.out_csv}") + + +if __name__ == "__main__": + main() diff --git a/experiments/exp2_timespan.py b/experiments/exp2_timespan.py new file mode 100644 index 0000000..5c866b2 --- /dev/null +++ b/experiments/exp2_timespan.py @@ -0,0 +1,317 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +Experiment 2: Effect of Timespan (ba-sir) + +- Observations: only TWO frames per run: + obs_time = { floor(T/2), T } + +- Sweep T from 3 to 10 (skip 1,2). +- For each T, regenerate a dataset file: + /exp2_timespan/T{T}/synthetic/ba-sir.pt + +- Compare: + 1) HERMES (hermes.py) + 2) CRI-MS (cri_ms.py) + 3) DHREC-MS (dhrec_ms.py) + +- Evaluation: + Use the existing inc/tester inside each script (they already do), + then this orchestrator reads the saved .pt result and writes a CSV + with f1 and nrmse (no plotting). + +Run from repo root: + python experiments/exp2_timespan.py --data_dir input --device_hermes cuda --device_dhrec cuda --device_cri cpu +""" + +from __future__ import annotations + +import os +import sys +import csv +import time +import argparse +import subprocess +from datetime import datetime +from typing import Dict, Any, List + +import torch +from torch_geometric.data import Data + +# --------------------------------------------------------------------- +# Make repo root importable (so `from inc.data import ...` works) +# --------------------------------------------------------------------- +ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) +if ROOT not in sys.path: + sys.path.insert(0, ROOT) + +from inc.data import data_load, data_make_states # type: ignore + + +CSV_FIELDS = [ + "timestamp", + "dataset", + "T", + "obs_time_mid", + "obs_time", + "method", + "device", + "seed", + "status", + "f1", + "nrmse", + "result_pt", + "data_dir_used", + "data_pt_used", + "cmd", + "elapsed_sec", +] + + +def _ensure_dir(path: str) -> None: + os.makedirs(path, exist_ok=True) + + +def _append_csv(path: str, row: Dict[str, Any]) -> None: + _ensure_dir(os.path.dirname(path)) + file_exists = os.path.exists(path) + with open(path, "a", newline="", encoding="utf-8") as f: + w = csv.DictWriter(f, fieldnames=CSV_FIELDS, extrasaction="ignore") + if not file_exists: + w.writeheader() + fixed_row = {k: row.get(k, "") for k in CSV_FIELDS} + w.writerow(fixed_row) + + +def _variant_timespan(base: Data, T_new: int) -> Data: + """ + Create a new Data with timespan clipped to T_new: + - tI, tR clipped to <= T_new+1 + - y rebuilt by data_make_states(T_new, tI_new, tR_new) + """ + assert T_new >= 1 + device = base.tI.device # typically CPU here + T_tensor = torch.tensor(T_new, dtype=base.T.dtype, device=device) + + cap = torch.full_like(base.tI, T_new + 1) + tI_new = torch.minimum(base.tI, cap) + + tR_new = None + if hasattr(base, "tR") and getattr(base, "tR") is not None: + capR = torch.full_like(base.tR, T_new + 1) + tR_new = torch.minimum(base.tR, capR) + + y_new = data_make_states(T_new, tI_new, tR_new) + + out = Data(edge_index=base.edge_index.clone(), y=y_new, T=T_tensor, tI=tI_new.clone()) + if tR_new is not None: + out.tR = tR_new.clone() + out.num_nodes = y_new.size(0) + return out + + +def _save_ba_sir_variant(data_dir_T: str, data_T: Data) -> str: + """ + Save variant dataset to: + /synthetic/ba-sir.pt + so that data_load('ba-sir', data_dir_T, device) will load this file. + """ + pt_path = os.path.join(data_dir_T, "synthetic", "ba-sir.pt") + _ensure_dir(os.path.dirname(pt_path)) + torch.save(data_T.cpu(), pt_path) + return pt_path + + +def _load_metrics_from_tester_pt(result_pt: str, dataset: str) -> Dict[str, float]: + """ + Each method script saves tester.res via torch.save(res, output). + res format: res[dataset]['f1'] = [..], res[dataset]['nrmse'] = [..] + """ + res = torch.load(result_pt, map_location="cpu", weights_only=False) + f1 = float(res[dataset]["f1"][-1]) + nrmse = float(res[dataset]["nrmse"][-1]) + return {"f1": f1, "nrmse": nrmse} + + +def _run_cmd(cmd: List[str]) -> float: + """ + Run a command and return wall-clock seconds. + """ + t0 = time.perf_counter() + subprocess.run(cmd, check=True) + return time.perf_counter() - t0 + + +def main() -> None: + p = argparse.ArgumentParser() + + # sweep setting + p.add_argument("--dataset", type=str, default="ba-sir") + p.add_argument("--seed", type=int, default=123456789) + p.add_argument("--data_dir", type=str, default="input") + p.add_argument("--T_min", type=int, default=3) + p.add_argument("--T_max", type=int, default=10) + + # where to store per-T datasets and per-run result .pt + p.add_argument("--gen_root", type=str, default="", help="default: /exp2_timespan") + p.add_argument("--out_root", type=str, default="output/exp2_timespan") + p.add_argument("--out_csv", type=str, default="output/exp2_timespan_metrics_ba_sir.csv") + p.add_argument("--overwrite_csv", action="store_true") + + # devices + p.add_argument("--device_hermes", type=str, default="cuda") + p.add_argument("--device_dhrec", type=str, default="cuda") + p.add_argument("--device_cri", type=str, default="cpu") + + # HERMES hyperparameters (keep fixed) + p.add_argument("--b_pI0", type=float, default=0.001) + p.add_argument("--b_pR0", type=float, default=0.001) + p.add_argument("--b_steps", type=int, default=500) + p.add_argument("--b_lr", type=float, default=0.003) + p.add_argument("--q_steps", type=int, default=250) + p.add_argument("--q_lr", type=float, default=0.003) + p.add_argument("--q_hid", type=int, default=16) + p.add_argument("--q_gnn", type=int, default=3) + p.add_argument("--q_mlp", type=int, default=2) + p.add_argument("--q_samples", type=int, default=10) + p.add_argument("--q_zlim", type=int, default=16) + p.add_argument("--p_coef", type=float, default=1.0) + p.add_argument("--t_samples", type=int, default=100) + p.add_argument("--t_steps", type=int, default=100) + p.add_argument("--t_keep", type=float, default=0.5) + + args = p.parse_args() + + if args.overwrite_csv and os.path.exists(args.out_csv): + os.remove(args.out_csv) + + gen_root = args.gen_root.strip() or os.path.join(args.data_dir, "exp2_timespan") + _ensure_dir(gen_root) + _ensure_dir(args.out_root) + + # Load base ba-sir once (CPU). For synthetic ba-sir, this is typically T=10 cached at /synthetic/ba-sir.pt. + base = data_load(args.dataset, args.data_dir, torch.device("cpu")) + base_T = int(base.T.item()) + + print(f"[load base] dataset={args.dataset} base_T={base_T} (expect >= {args.T_max})") + + py = sys.executable + + # sweep + for T in range(args.T_min, args.T_max + 1): + if T <= 2: + continue + if T > base_T: + print(f"[skip] T={T} > base_T={base_T}") + continue + + obs_mid = T // 2 + obs_time = [obs_mid, T] # exactly two frames + obs_time_str = ",".join(map(str, obs_time)) + + # 1) regenerate data for this T into a dedicated data_dir + data_dir_T = os.path.join(gen_root, f"T{T}") + data_T = _variant_timespan(base, T) + data_pt = _save_ba_sir_variant(data_dir_T, data_T) + print(f"[data] T={T} saved -> {data_pt} (obs={obs_time_str})") + + # 2) run methods + # 2.1 HERMES: obs_time arg is "extra observed times"; final T always included inside hermes.py + hermes_out = os.path.join(args.out_root, f"hermes_T{T}.pt") + cmd_hermes = [ + py, "hermes.py", + "--dataset", args.dataset, + "--seed", str(args.seed), + "--data_dir", data_dir_T, + "--output", hermes_out, + "--device", args.device_hermes, + "--obs_time", str(obs_mid), + "--b_pI0", str(args.b_pI0), "--b_pR0", str(args.b_pR0), "--b_steps", str(args.b_steps), "--b_lr", str(args.b_lr), + "--q_steps", str(args.q_steps), "--q_lr", str(args.q_lr), "--q_hid", str(args.q_hid), "--q_gnn", str(args.q_gnn), + "--q_mlp", str(args.q_mlp), "--q_samples", str(args.q_samples), "--q_zlim", str(args.q_zlim), + "--p_coef", str(args.p_coef), + "--t_samples", str(args.t_samples), "--t_steps", str(args.t_steps), "--t_keep", str(args.t_keep), + ] + + # 2.2 CRI-MS: pass BOTH mid and T explicitly (cri_ms.py doesn't auto-append T) + cri_out = os.path.join(args.out_root, f"cri_T{T}.pt") + cmd_cri = [ + py, "cri_ms.py", + "--dataset", args.dataset, + "--seed", str(args.seed), + "--data_dir", data_dir_T, + "--output", cri_out, + "--device", args.device_cri, + "--obs_ts", obs_time_str, + ] + + # 2.3 DHREC-MS: pass BOTH mid and T (dhrec_ms.py will ensure T included anyway) + dhrec_out = os.path.join(args.out_root, f"dhrec_T{T}.pt") + cmd_dhrec = [ + py, "dhrec_ms.py", + "--dataset", args.dataset, + "--seed", str(args.seed), + "--data_dir", data_dir_T, + "--output", dhrec_out, + "--device", args.device_dhrec, + "--b_pI0", str(args.b_pI0), "--b_pR0", str(args.b_pR0), "--b_steps", str(args.b_steps), "--b_lr", str(args.b_lr), + "--obs_ts", obs_time_str, + ] + + runs = [ + ("hermes", args.device_hermes, cmd_hermes, hermes_out), + ("cri", args.device_cri, cmd_cri, cri_out), + ("dhrec", args.device_dhrec, cmd_dhrec, dhrec_out), + ] + + for method, device, cmd, out_pt in runs: + row = { + "timestamp": datetime.now().isoformat(timespec="seconds"), + "dataset": args.dataset, + "T": T, + "obs_time_mid": obs_mid, + "obs_time": obs_time_str, + "method": method, + "device": device, + "seed": args.seed, + "data_dir_used": data_dir_T, + "data_pt_used": data_pt, + "result_pt": out_pt, + "cmd": " ".join(cmd), + } + + status = "ok" + elapsed = float("nan") + f1 = float("nan") + nrmse = float("nan") + + try: + elapsed = _run_cmd(cmd) + mets = _load_metrics_from_tester_pt(out_pt, args.dataset) + f1, nrmse = mets["f1"], mets["nrmse"] + except subprocess.CalledProcessError: + status = "error" + except FileNotFoundError: + status = "missing_output" + except Exception: + status = "error" + + row.update({ + "status": status, + "elapsed_sec": elapsed, + "f1": f1, + "nrmse": nrmse, + }) + _append_csv(args.out_csv, row) + + print(f"[done] T={T} method={method:6s} status={status} " + f"f1={f1 if f1==f1 else float('nan'):.4f} " + f"nrmse={nrmse if nrmse==nrmse else float('nan'):.4f} " + f"sec={elapsed if elapsed==elapsed else float('nan'):.2f}") + + print(f"[ok] wrote CSV -> {args.out_csv}") + + +if __name__ == "__main__": + main() diff --git a/hermes.py b/hermes.py index febe412..3f395bb 100644 --- a/hermes.py +++ b/hermes.py @@ -346,7 +346,12 @@ def main(data): y_pred = y_pred[:, : data.T.item()].cummax(dim = 1).values return y_pred -args = get_args() -tester = Tester(args.data_dir, args.device, main) -tester.test([args.dataset], seed = args.seed, rep = 1) -tester.save(args.output) \ No newline at end of file +def cli_main(): + global args + args = get_args() + tester = Tester(args.data_dir, args.device, main) + tester.test([args.dataset], seed=args.seed, rep=1) + tester.save(args.output) + +if __name__ == "__main__": + cli_main() diff --git a/input/exp1_scalability/ba-si-n1000-T1-g0-d1.pt b/input/exp1_scalability/ba-si-n1000-T1-g0-d1.pt new file mode 100644 index 0000000000000000000000000000000000000000..f3d68529109e7d812f3f5978fc7ffba74ecc81b4 GIT binary patch literal 178271 zcmeFa2i(tP-^Xnwl&nx>%ics}?^*Uk3`Ek`Bj7}2L+$8PPqcJ9-w zb40)H9Si-P9{GQgRgL^V;qmg04v#-PJVBo3w_U2FUFS|+JGTpqi0Id@W6%EK38Tnf z?UHw;OU}u{6E#m6nLVO&@BV%I{l5j=6h#4z|MtIJMB?U&Ba7&Nsl}bcZ~osiMka=J z?VPMe&psW(dNx9ZC%N)W?`@2=su3=h9GPxEA#On*OC_uid4 z4-QX$x#U9a21Io4**{skew{lE=-#uFO&?e+S@>;v@-)BsQZ8x>PjNX*vbueGcl}SD zDVxW?Q`PPrBa(%u%iAzKeMI;j4a4tjp465`zJTxy-7dAGn{`+3 z*(WUGKb2?vyYiypcXdmBsg8*7OqVL9vdnejN7k8Ta<|B$vqpqxYZ#vW@1lDT=+&-w zpH7|ohv%sGX~MrtcqlyQN@cw!V6u>oj330422`Yi!=-`+We+VPuI0ySl@2p z#hNF$^#6AV3orinEAQ03OP2xt!%O^a?cY88p1ugFmp*0Ec7$vTPtej9yjwTtLC zpmTVs8Wr9t{`cQybHmH#duv=O^Y6dE%c)f)ymXyp1O9et*Eg(R=iU+FWxB;5&|yG_ zLiOw1oaFU=kJL|p>CHv{r+&z3mNQ=Z|9zSj>e8oQuds-A13UNYf9a)Gubd`c%m4XT zG_t16U!NQsIaDJ56S}6{@$!YJiBZ@ttY`PGy?b4Hj_Q?j{cn~1&+1#2ED_m&3WLHw zi|p{o|AcNp=6J>a_Xc!$D57(p$RS(3a-nF}Yg2l5kBI2mIkFAo%bmS5a(YMpCv+P! z#w+^2w;^)DDqN&`<$^J5z@^IK{^EKd${vWEnvs8T{NMW_a^;Bpi`?lV|Kb?e4_9Iw zTncXg{)=N=KScJx)fxwpp=ab@9Ivl_h>Qj!|Kb?e5C50PLF8pZxV1b50Mv{k$=}KBd>%b|Kj-X^usNetBv>H&Rgk+I9~3E*!Dwm`Z11|JK{?B zL*%Ey$iFz=SpAUVa&zOw@p4B*exr{3`@hOmmme?Q)naPOzh0Pz_S`7Eo%Y7@az}*f zhqUz5m1a81{cmM@oif4?fcmc$3sT?p!a}s?-(q3vzt${5z1JIy(ypt`V$^e` zSe){1G?t*R*>WMDz#Xrc>*n=#C{c7<(%D>VqOSv~1%hB%GvOM{( zH7ii>jmnC&|3+ga+8t9?CSPn>h5R=vtJ40f#cGrv6|3VHQ`R8gm1a%K{im!&Ik9DJ z^8ZuTp`3rpx|DORSdaRm%lq-a-uM9R`nOo0`ma|upnX9$#D1mNi1K2}#^ejdCZzvU zHl>`HvKjfV6q{4twPFkEyH;#Teg7?N#q(Y%wx&GU27RNjE$#ic*pB+6%l7z(Vi@T$ zWe4&_#g6#Jl%2_Uwb+I7qheS5V#{vikBZ&#i!OWM9~B?OFShJS{@Ah?`D4r8HijIQ-6?;V!zUy zL%E?im-K7RdDI(SK8FAG!uhl($S1JJluwc`$OYJA%7x_nr+kWXt~D1?Z*;jBf4Kx5 zT`t8xrd&q8sJI-zf14|4M^t90h=S4qFxe2scS@paPW8|c{b z0Qp1lP13J4-=f?g-^LzW9wfg!gpMu`_L8jJrqxq9^^;ZgFJ&h6wi__KSoEzbNKyRJWu`d6ZEH1@H5io=jb56 zz#bJZ;P>S<;#brozedND-;gi5{1*QpFJk`=ejnfuX!#>Lw)~0wLH>+g{(_D!f5kt# z{0)EkJ318qApKAHSAa47FP4rU|8fB_WdicaglL%vEpI}HVq(&7jwO?jFBFrK9$O|O ze{^{Z{QpCLkHO&y9`6i z4rtjC-6@*vj9+Zoh5WKBI=bwJe^l&_Uk~_TfIZQ&H#!vikRBDo@e8sa_UN)d{y|1y zmjloP;h+Eqqva5EkVCP{Vd$80IQd2dI0`LCqk|lSJ;<@x1$V4@8r1 z;uqvw*xwGpgV^OE^g96_K?ivhyF7+|FTnTFCqnQf_79`rY0{(ONBGGz=(8dCG4>$O zVLuN)fuDxpXV^av!7s31h$X)wpZppfGgI>?gPOTl|Xur&7j0xXZN1S?0ARq&J5(XuAG7OWj$9duoIKWqRS!p89L|K-|C zQ?(_!RS34m-Zl!hCp|0#J7Dh^f}OB;39u_#c0&i*9lPv-4zefqAbVl&9bg}HKiD4* z2=F2F5I8izVQ4uVEk~l|D0Gmcv5$dc100W@5a2|#oP?ed;8e7nj+Qgfj|4aiEoYQnXx#mQSODdG!5zs(`I1^bHuzJv~PFZNg9>jA!rmIu-D5c+U{@1W%o^w9v1q2+P(djY=A4e*rJRui)44VhDbR{YUr{{5cB#O8Or#MZ(MRS}3L> zU8Y8-2{0`>U4VC@Wd^j&gw71Jz^pJk%nysgVz3M>2dlvvuokQX>jqd4-553tur<0( zfbGzI;7~XWj)jxqBXAy^4;R3N@L9MTu7#W7R`?R!1qsXJ@ml}$znh%s^6UN2{~qbh z=v!fOcpFR$?}L?LRagVoge_r5*cT3m6W}B`4bFj2K!1np?<4&kp})iPcWwTj%HP@e zyNt?oKvUQPhQV=g67)XHbEfCSo9MvA@OGF9W`{YT@opJ(LukCz7wtZ{8@&g<3s1si zywBUA<7qs4Dx3}%!)M?|xEsC+kHMSh@5IpYkOZ9$W`bE^PM8nohmMm9=*qAWYyrE# zt}p^lgp=V+_!xW+ZiD{*&EE<6`yI!@t-RmdunMdL@1>vP(;uVY3F!Qas-Ke5o@DSA z=)6jSP6g9I`!^dp2P^`M!xFFxtOXlF`!x)0|8_(7fcEoX^f2iBTZ(=fZh||Y^YA0I zzw`3{^rVc7+|d1}G`co)ADDt(2H%Il@#Hvh-F3V;PHv~aU60eFGr(*x2fQ2R zh6SPHs3f{9EC=hr2cYAtDY_#Z2xq|Aa4wt&owtvp7sB0eA3OjL!H?jVF!+18^o)m! z(Dl9++WB1{?Yi&x!hYBvf_^U>iS~QSV`#qz>_NW@e}dkN-N`s{Unqty0ZYMhuoLuL zY20s|{VM6kwG|mZ?O`YAe()OFaTaP_53W~^yOi{IR#*rYhOY0`(eC^Aqnp4E(0zXd zdJG&3CqUx?$I~+G%i%6~8vYIw)BkB;2ACffhxfpiuoqkg{Z8V2e=_<#4NM0!!UfQC z;3o7f^nVf90gi#bC->LjdkF3u$$3uSi|ar-?5+bD(a!rk=-~cb0DDn*4=e>s!-}vH ztOjeq+R*jF_udM78yE(=!y(XpX9?Q<<{9)lXdJi&y%W9<{XP57=w!Ul6fiY(Tx3RP zgE?U{*bcUb-QWP|Ixz6_L-&9);auqboa?CjpZ7q88Bd6W`~}8Twh&JtCH@%SQqWO*bePDGJYD0eJY#|7r`aa^>H6&z&x-BtOjd9?`3KcmSS)@pzuQq3c>DwBxQOx)GcPAAz&sBDfeD*KJ1cf_~rkyZblT{jQyg z{`7ol|GFM|FP((`^&I#pdM6;W@HG*FX1J&kdzWp9mj@Q=#kIEj({Zm=0!xjo|?J z1oR%z`Rlx$Pk(t2{1*BHXuo++V1H%cc^x--(azs{&<$Y+X#e#=PlL1Iqwoo69DWl0 zJ^U5^1q;$o)nRMs`Zf;jz0gwhTDTj&4!wshP5;@iP0)_N=h1$jIg9pQry>35efl=E z@r3g+3(t|&=Yf@=`+g0y@n&;$80-YQz#gz4G%gr|Ha-}Mo&sk=_kZ{O)7b6*jPzFp z=(^e+?YYMOcYIGFeFk)WU5$2KJ&OJUCZK;a!+h|5*Z?+zqoC&#zk_;D_afy4eivdl`dX1&z1XqThy&`}jP!@l$qmJ81kf7X1u-2j<~< zdcl$K5jY<{1N{!@{bMtpe*&Bc{k~*e5?m+J^BnnML0AgbgtcL7*bzD&H>2NyXW*C6 z_|kFcy4sWHe-ti&&iiHPKj2@`{_+0hPM+8K*@}91z#m{~@>{R(VGVZUg9-HWyU_1y zbt$JFbY8bbcY}SQ^YR6ThU=hl)w}4Vw97cs zdDa}e`%Hhd=NR{wdDxwg#+6&KpMahhlkj{kp}$Ld6nzf<00;5>&Xc0FyEH5ZD?#^f z=XrbV?hk{}Bj7k_|GbAj4L^fF!<%@XT(BZ+3q3zSgf`xGKE~&{yhoXhUIe$n0z9|( z&_mI);9lrD(vJ6T=-4DV^55})gVqXL|K;x+c=&xZs+FJ)Ufv%tKV+*l+o=iylbHF07 zA{+`|fG@#U;jKL9Qka?cxbNjgyM7i!yPw~SZUCFX7H~6s5xxqKz%wub&vi4*1YIYq zqwj~k;dto%_`7J||4(Sokp*~e=XnO6+w+mX!!h1FMY{LG^J(v1=sfI4J=H1SIMUxu z9HJciuQd6qLHGIIX!n(+=(nKr-TSdo)a!X>6Wa4kcG_c{+#G!bW}#i(;Yzq4R-oP< zus@s*_rbE%TNV19@I%TegI{@A5jKPm!bhO<=OFrH=sDtP>Nx<%@I1y5u2Z$BzZ;wb zjSD(cuItkl^km9i1&yl)Q_lu?3i`XWyD8s&*LigwdvIThpA}CT`(`K3tiXTmsVh3 z4L^bfX{T}E9&|41DG6O?_M?x&0+i!9$@h5-dpr8Wea_$Qc`xGq+APXj23@x-e*)$C zyO062(|DjR`PxIr`5^RYI0YIPeTsHpA4omZ;4bLBc72}5^GinBF&w@M-N&j>em^(@ zI{y};*TVI12RsIC-!R%U5*j!9yM$+H-|sLU?I{VLhTdno?z(=s4+Z1cCbYvi!}ZDc z^a<&X$HBDE^Wi%5>$Iyl{v}{-ct7k7ABC<{t{dC2zYgDp=io*7JNyHtryXUW`&?sm z4>%e+?i>%U|Jz7U!Sm#X1z{oB9=e`+-*3Eop7hMT7r)y+hyDaErTr)1Nm!0{42Ayg zxC;4<8=B+)5Om%1y*Uqzf6h|gx3DPfjDUW(sYQ7$VLLb#&Va@jr6_L%d~_wxmmzXy76@eMi~<+p*8;H&T)%*Aunf@7)I-(NaUtiLSf4}h*)7tkdr zw>~s}@Vow6>NyIHk8)7&KIr=9zFw8`C&M@4uh8*+3-!8B6i2&1c&_vJ5T{8Gp2wEa zzV*=ksxkHXzTMx)V&4bP!tdZ?)bjz{PPqwrUe_u6e+TxX(0$VLT^q{vJ&ZxS&VP>n z7AB)!*Ma$g_I!n|O?`fUzK{IIX_e8B!6k49%tASizYo#wyWSTWN8C$4RDk}TBOmSX z_w4?D*mINdt>eS%Jq4S{(`M-nddA@S63jS4Leb@(%hQ9v|==b1{(0Jr8 z^n*Otv+xi+2|tFvLZ3T~c8-SrzH2h=^>;MJ$K9##Em)EAcEBGgCn@>;-Hy*`TsDgQ z^P&5&{pmRKe%kj|iSi$U+hAYn8x5C1_xbVEGY>9@U%_vo_ZzKv9`7&Kqm4K6Qm^rv z_vFSm?t9+fe~Eu$+GX7DJ@B*G_d?GB#%D zyr)@={cTv2_VC+k?~n4{EFdM54zqhCtq3mDFXUC zFxLgwHP0DCDR(3s3m=B|=OXkAa2tFRI!;_)a?`#d(Eja$_IW+`=cN7a*Tc|Vc#cuj zV|!jhSED?~Wmoj$urB3HfvqTSAUptHp}gD3x5oUWzX0pu=Qu3{JXC=r60U! zGmdk984oWa|4R5e{2khU_up42*ZHxZ=Xo8v?#-aRe*aEIzN)Y;bX|3vI6jS+3(=0I zuo-lGy1p&N?)R13s6P+%JYhVN6o1d_o6w)ZAED>{&OBdN_z-*+u7U5sGw^%Zgz_GN z#@%0|8`I8Fa2nhNXV70~;cUvae$QLAD6bdn13lNSKp%&md;A_}yfuk>H^En+-{sd+ z|1s#f+XUyBeqogBx@BB=D}I^a zaQFlC-gPPEJVH6<&yHVhxE1cFo!>zB8RL$Xfqn}kpZkMxt>f2o#?7?j9Qn^f=S4iq zDGSTP4sbvG9D3h;H|;DBAEh48$?ls|sBb=e2D%P7Uv^XeLFj#t-&yKW&wbR>3r>Qr z4?QTqC-mNPIC>1U9ox}=LF2`?v|~El26sZ&>5J&UU_9EH61uLmM-PV!q4zGH`@Dbf z_rsldzOL{=*pBwjg(<23PH6md0qr<2UTjM_6QJYN^MP?tcglCYH4d7F{WNqu{f2&m zdR<@L4;-hiKgQ#a(;uGKv(YYppXGb^{@(Rx6@DA2cPDf|bA3F8-TkI8?Q02#z!`8e zdP;VP(|4gTzMKBEiPvK9{@#nhixcZ)YmXP0l!SA^T$^QrRcd_r$4&%<7 zDW?te_chb-^W5ux>HJSexxSy5(GAFVly(KbPkDayp85>+dd{g$Io`j!pZJ}?-=W!m zL&@)X!h2ZHXP?uKFQD(ydF%S&eiD!NJFeD|uP^Ny0{xD74*#+E8TbB){*(4KC7*Gb z{bRi4y7)5Xoq_u(rxWGug@^6BkGhDY%kiRqe4?*9*_XUaY%LI)} zv*AAi&Vf^Czw4^|ljop?wC8!aG0;x?=LhUR!Mmy7b?aVqRp>s^8ExEP{OtO?k97AD z;{d-;HRCyY!h}3eURWR6PsRh@vpr6Gz5g=~b00ZGy8G-mX!n91$8@M0uY;&)xYc|4FzI zdYnOci}gY_PAaS!q2#1B<21I7m#lXe#5Yzf%a!T(v8bppW5Q*e(wG5 zqqHLj`E19X_?LjKq4DF}WH&Jp{dvS&v^j?C0TD+W9Foo=ZVJ_rq4u zd(v*`yC}CMd>Gn}0o3FDyX)l#*c(yKBjnovpP)U)O|HxC`^H06$nSe|oKK`2&yOvs z?=i}Cf4BZ_q&v>f;^+C%`05$zdk(IH8L2NHd<`a}+@i1%v^~z_@|5E~d7Av`@ox;r zL&v%E%=u*eSdQ{~!Eoq(TN29m-h3c>E8Gn~fC;GAI3hLL^}P(`yHAWF-Fuppr0;_6 z>-JNA{5{t!rT-`5e>drCv404Ce4pYg8y%YOWiLC-@q$sYkd2hTwt zg~p@qTVrUi>!NX+_pxu$4%f#flye#e-^Ur+<$G{_?1KMrn1*@>L-*_N(G93~qy0d> zn@N8UCMVr<(=z;yLFdE!wCgu$+?oMD&%dtQ?$>XTZXELz^%;LwBmD&Jao)X5J%>qm z9rGUEet4Ahvg99wwqE0@;-n9QTcG>jOxjhJ{Nv#~Xus{HUB-!ZNgqRg*9*@(X-RiI z7>Q0u{V8D?@_U{%p4fr^1pHjD-=iHpNca5x7uxwBPQH22b+!uixqo|3Grslr1Kw}s zq1@@veaU<25tO%=c6krzzLtf2d0}2&DIPNFo?>Qqc&vOU$dXH;7&=mWA z^0_`be{+$41bhX$F7+h8@s#VW@w)564)WE&-?-oNQZwrJJ*C3W^~U-8Ipugx`JQ~s z@!tY7V)vft9`Yr^?s)kKU7z%Y@G$-AIBkVJ40;cc8SVHPgmzsmMSGmj#!1eXcc|a> zeggjfUc!0mx@w$zihBGmQJ8XSK=1S2Hxpxbo_UVG8+!xT5c)n=p)W$$2k$MvrX8Mt ztK!!QE`km4yMUgJE{0wLeNNXM*K5bG_h9Aer%mKjKZyM}^gSlW?>jh)`hDLk(LPU3 z^5?T2cmz7Wj1T%?zreWh`MXiylkjDD7xnFgiO6q%55=An|JiUpY(x6T@B(yRRK%|~ zG%gv9o&i0-`F-O{%2`Z%&cN09S0lYHYzBXWEAcOazj2V~M%!yVXx!2q|3c(H25-ea z0nUZS6Z5IJKmHHHHE;pxbI@ban<=LK z?|AW?`xEK$DL*On-obd`9QI&bzX*E*+EE&|pxi^G7sNgoK1ljTXdL+!{ouN8Jdy!F z@1@hCEpHg*J_3J)=Fft^=aU-f+R*jvTeSV*y6$<*dOW8jp`K)LD|ElPm*+PQ@t&t> zz`r;3roew6Y>waCq`N=5zPS!fC*67MygH8Eb!-FWZG-3GE%@btAHgoP!}H)Q{9Olr zLffClLG`eI3+)HX8&0|IYro)Ez8%!*HX@IcoO>@^e1o<&pQ`7?{CL{5PT1Q z3jNMw|GbFZe)V3*{m199zBlk+LpwZQWyRhYu7|hcmk#=U!~3z8_`L~@kGv;(9=rQO zb+qr<^Xo?K^hbZ{^Lx7MR%`5C;k$4w`5vSk<97T{K=-psls_HTrySo;W%N3D3x12yyWnwn59K$6*82eY-h~nPKLo#k z7sz)MzxSc>cxTdG=Z$~85BX0!$Tt-K^=RWY?|od49>#AL+yI?N4QSs1>3v~c{EXl2&nftChwk&s$Txs|o8VrUi1gIZ z`_}`cKTdhxCw+jfj=%52IOBKh#_O(Q-LTuQuc3ER-XVAqyX&~??GxCITRud)54+y@ zo*Y*z@G~Cvd3IBu`^lH++wr%(EwGP)o-5zN&+^LBU*Y&|g0H|6@GE$p{N6te#qN2~ ze$7a^wV-iqdh#EEcarXY?YOa@N|5e*d4_zyLg#-I^i=5h*^BN#ImTUw(SOkX5%{&j z?tbQcaX-p~UpD*-!p@ZQG&DXuhyMipoL7yo@5i1CZTvn3{VL_zAIH$n_jTld9WEi= zdH5dXWFUPZ^qk}R@A=ztJDPls7vuG#*o{j%p!ZPjOzJl-IE!w9pXbyE@qeH6GtjtT zEdIVH&)Wm>n@IXil;eDG-!<;~9RC9N55eE@ZN4Az^W5+f{_)@d_!M+qITvWp=cNA( zyHlU$nPz_q61XfSwoAQJ(kLJJ4^yS@=7zT$d7JUkr^0eJ=^HTaN3` z57;wOzR&CUd=k6Qe}w)wo;Hqj-7HCY=kY6z?n^nH;fJK(M!Nkq4()j&AKLZlGxRFx zzSg)zX!jD&cn5o_bRMM`ef+3^*-8tz8W6DSPlQ1V0`S}qhG+@0DB^|>$ClD9AjKE9lvkkHuAe}F2z0%yYcsM z?2BM+=ss}-Jpn!fM^R1>=)8QD^cv84;J7@5-F2^$?SYox7~Kf}WN6n3U&QYH zoa>0?Z^F-U}7&;z&AII@K1s$ig@b3cMFKl0N?6&_k^eO0mRRP@wdJZmy zHlOdcBlZ=vXEk(Qts;FFd>eiQ6Tvsh=RA4~ot*SZ&~fwu>Gxni2v@@b@FITM(XK1* zuRE|C58Q$t4VS{7C}%HpT`~S~e$~N$2y~qHL(hcYQI74lohwKm2^|N|pk1Gw2gWn5 z7p{Aa@Q;9fV1GCQy05qn--X@zY`H#H2hv@4?f-Px9S>{KAH$nr3Fv&Di~gDZIEP*a zf5+}Rm<+r9mJ)5hSpFT@oyY6Y&%#2Ghd( zU}abp)_^r(OV|}m@9ce!=bOZId(}{B~a8FgPAM zFFc1Y#_oLEh<1LxiFVvL4ieL!j)x@ZbTAXl3Uk7IFh6vhR6tjTjbIDd1$KoIa3Y)x z{da%L4uZcYtHOJ(1Mj83cFgGj+ z9Y-b6Wnnp32R;BDUro^+;XpV8&W3a0Jm|cA9K8_khW__B96%p}AHgr-Z}4}Rp7Br- zy5849JHP9rUH5yS`@x6c5I7Rff{#JJ%j`k>9oBm~zbD+uIB{PnhAsh1!E(^^Zy&Vx zB3sa&i#*R)Wc;*;ouK=H=Mu+RsC7NKUODbk(!W_@Ay^o?zE?-P@86GZ0y{wW{SoLf za4ehvjRzc0%djtpyWnZ)I7rOQvp4?x9?;*HvBEpo9B!0qjNLJ+Krk4J*P* zuo|oZYeUxy-+L?UZD1Je4u`;Da0zt3c?P`>J`cCRo$!6=@8Ew%C*%F4fT^M5A~QN0 z%n6&pcF^D1bVK{QJJ*RZ*j*pyqo0He;S%^9+zwqwj-x+-elPkMoyc(oi@=i5ea-cw zKX%8Z_aEL%c#rM<)kgf^hlv@d`C%bg6NbZia0y%kJ*T{bc6>Xo9nY>esd;aALdS6? zbY_?Z=75fK_czCRaG!I3t3^J?tLv5fTz%5JLdR=Qv~iQ`**NUuq3hQIwC9%HXyc|= z(Qm;E@GJNmOvw9B3p2rd(7389+VR>6?f&%y+VA-L(Ffp{(D=&n=K7bGapAg`6`d0n zhwf)B(Cwk)vm3ex^d5IE+WTMEQTM-tq!(s9wS%6IJ(qiq_WbVnJC48iFUja%*Hg!z zah2=pUHD~%*`eni*H_oms-(Lw)hQknLccS0 zM0bG^a1fjW=fNe=?-~2i$DtP@X&4`^pyMkH-2-~x<34U&@Asci@N<85JULzp(ci^j zMffD#43%F0{Z*uc09*exC@q}9aUiq=<~S_JMQnGUf0=d=a1A^FPr-OR*WJ){trFUCR}DkD42|nHqjy1n2lzGm8|d%M zQqiBDFYRB~qeG-8p?^IGK8p4}_i?oM3GRO>cz(x=@4o|Fig7B61ub%vO?EC_gT*krAVI$ABIz*>)S0nZ%UXBW`m94 z0Qdy-9?<#gyq!;fJqzE0A3*!fdjk6_1JCQY$%}UW-h*xkJ3#xd4|*D$1s{b^K;!U} z=zL)W))XzzuVqSwOR@O9`tY-#$>erolYv{l2jc zZ9L(8%))bI^?6_==)PYAZM@kW9R@qWF0cpe2aOAcpp6ekqNl)_(EZqL5j;QST1;11wE`>-9aX!ESsof&P6LzC`(T zDW@KEUbjVegMFd%@*A|{VJG#y1mA$aLC-<2QU52fDeX85Gf-|u=>9($y%@d-9k&B2J8p@{^~efgx&FGoRgh)7+<+A#Q zM$q@s2krVf3_TIfh0Eb?=s7zF{Za&$gpQjzXvd}J=92WwebD{K{oV7(_oP>(Uq-^& z(0$qGDo=ftq4!3KDJL~_z5NA!6XoTC#s%(Y?RjqRJ9?oVzk|@T;W}ts^)5Oo?J|yZ zo;AnrKGPrVImZ2E9(L!WaphL*C!puWBs^bB=1MW2H|z(G8}^Q0*4E)C1UO3?k= zdEOqo`@>-L2sjSfKkuPW!_VN)@Ft!o7pw@|LeI|+p^dklkMVgf?@?x>7r|}N-~D(G zJ(PZ!1@}VNk*2h-C-nKIqF2C`a1-=-KScis@1_0D8~48#u)heugYjs$-%AFf-S2-x z-%PuWe@~-xP;c;@?S2qWdN6){68j>!0UA#oKz|M6(cU_+33UB*A6tms^JGHWp92OyI@{e7rL&wFRj47 z8h!-*9iDOE9{MjA^^}CJGyBm;VFAkVoaFmFhP@sA;Xdc@FTEG>er*=zErYIGmOp{= z{9VWZ+G#vcmwfG^<9rZ$G@JsBi#|oWuMec2X>b?xUb{Ze_IW;Bhkl**7015>tPSsnz2T$Kb;@;PJNDP%yYL*m2!Drv!1T1E40NAs zjP3zPL&u%t!S#O|=_z=g+^`@l1lvQ`Gw=J2ch8fanfKy%+vm`qz@@bR1Uw1L(T<_; z0_9a9pK(KT{2zj@d%id4f$`5-%KH`;rQ8VUcbi(2*AljaQ{fC~d{K(>M!;8~_W;Jb zu1ELt9RBXn=W#z@K>2&1_ZHuvvr&E_=DStA26aESv@3&B|`$TcH`-A5?e-Ck*^x%1H z8SPsS-LD!`pYPlKeJu8U@GSfeK1Mws!0nWqkmq%svj6?P^ik5?Cq3V_p+1*SLoW*=lAFP$Zwoh8T}Ys0(ZbHl;imO5beI}eUWj*z4Su`=K@-w%6kGQM?ucwREjT}*pU!8>S&@q+WwbHD3hXX^9!ArFv$0(87ThxYtYgYvyE z*vj+xyDEPNZF?J#zY7e9?%Si#{XSkeuhu14}^n zb-&vf4|SxT!EiL31J}ThVM6-F`={+_&kZy6qg`oXceoMO!LJ!K&KZK<1rNc)(Dfh< z^?HtV-SwWy-@UdZUtef^b`SaY(VkB9lkcZD_MPw``L~cS1NOGqoww7_@$pLut#2dx zedxWC_bKBjcRDnFcD_2!9KXA$?>KZmlp+6jFg?##4pzaxDy$Ftz|qk6zXAOo{1F8=?RgPCMR|MhGd?mt>x5r1{OUp1yXE97OFuq~ChR|ML>UC=(S=l-0u-~DOejL`NoGGvsIt;a36GCx{G#Kg>9khlI!no$_o#a^98yr`N~1p^_ghbpMv;TfsTLo^``WL z_iV;-&M)KPMdV)zUx&X#+wcDS3gtRK_VYZiL)X0-wAb(7smNCq)`hOCjuXeH@p2*B z(G)g=j!)OOrP%$xavSyMfu1LfN0Q?2d3_W5Gx#I)yx*DU>k1!&&%!nE9e4(Q51UZl zBha|}Yjk7UISNjLyWkA^>nxm2xz_J_s}|+;f_5$|($u_nV_z!YAN5c-C~D-*dn?wD;&Qq9@U=AE5UetI)=S zb7|Kxcpvo~hLtF<9()eEUi6~8n&fw!I*y*9oZHCfdggxTy+eKSufZ>ja$UEK3vb0Q z6C4hIfZn?%C-dDsE&ho3|5oA0Kb<>8~$<2l)Va|-p%htELQ0q4tZ%0CFb@9{fJJ?goSdV0Z0 z(Dk7Q<@bc%dk#mBfwp5i`Y&j_*p_xohuh#z=sJB7{TGZ!J5xf}mGG#dDwc z5B`3*6VKNbJ_y^<-nlR(_1_7Ne=eXM2gZwSDQ5z7oO(Vm4(d+%uD8ZP)3Be0j;G(y zPf)MxtNVfD)b+=B{Bioj^LjSgt`Hs@A;P)xdkKR+Cp-p?++VKVSJvwh)Kip5^(SFC(I`Z|UT|=PX5zpa27C+tzdY5)9y{U>-g^}BA}i>?aYCpx2z8;qY_pZAgO zK4Ki;_o-$)M^BiL=gAA}L;J~iz}+Kjl9O z7edc7zMu4z?|I7aM62ls_u&NOZvo$cY4I}-ONV_Hblv|Ior?N8K<@*mQqBtK{^z+Q zHTmn1&-1AJ=SSr8KJ{Jt!+ZJ6$Mx%8^6i4Z;9mki z`~3mhy&0y&e=B?gy1%6-|61s}Ig0YBDhIR1I`F?bv0yT6>Ie4od-YA}9}LFZ{Mp4Wcphkh8k z9`&QXC&+Icf0KUDed{j#Cej|)%R%@V7mTFbKj8xMO~G#%_A}7_tVgpur)M(e4Bg;$(ITagvNmtXt(jf2jn|RdCt!_u@4}>-y^<5dw%>G z{R^x@InD>m4X0ku(+4QO0Qok;lhAoJneq;he=q6ZLhlE<;5QyFhsL2lqL;!w)bIM^ z`c^4Wer5chgYLgMXpj4;acgh#`@Osg`Mfu-PQHhr_c817ONaeD+)6t?g~oF!sONsz z3VKi44Sg5owuBEu+cAK8ynlDS`~Z6+%6Wu*8{iYP$GFLL*?r%5s0#UgZ;tbcl;ioa zCG|Z>Lw(P|buc6K<%6%mWRzPJHiEXtd0d`y+$T?yKRy19 z;dtmccb++)j33KUUN0C9y>ClG`QDokL~n(=;Ri4Q^%_T{M!UY3p?vp=QKWlMlalmZ z(0$#0%8$S2nx*vrMEvh2eJ%D6q3>@5eqY1SXm?`rEyv$=-u0y*em?(GXyX(6{RPU| zLpjHx^Tm6ktkj-b<9`w!q@E)%2l<_sS@AR8b${89|1s!!s3!R%py%K@ z=%dhh)O~9V?R8x=Zu36&E!yGw*o1OU!{GZkL%Vzru8&>t9}d${?_lVD{XMz?^=`Bu z$agd8@4@7xdv02W-!bTXc%OFt28~-Y;OF_*b=&>=Ez*r+o}xbE&uXNfpgqpJm#OD4 z>8@km!`ly!l3telW6;)XJXM_ZVQ>p{|C>p>%94LPoCoc-BrIqX+4pzyCrz|HH{Q54z4)p+5I-&uPZD{(ivw zjXab)9l9@hFFk_t_R=oz0o~WKkS{N+4jsqF$(~EUq#VcnWc)p6s z^KVuBI>AM-0e%U$Et4DX`8oiGvk?eC%3bK*Z6&WCMC{}^6?&Wnop z)rQ6;qtP><=QqD^d`UTrY0nwB8vkmf*M-gCZ*V34MesKc^4w^9jR%cen&V%H{Kw#} z*eAfb(0F1#_4ddAVYmh^Abk#cEP6BLR7TGx{X^(E*>%f!=so_&k9KU~*6k6DlBlqA%X3~q()H}~@V#v$JG z6b<SOe_HUv6V0ptS*M035{0dkv>NBj0E6 zB6NODqP#uO=V*stZtUJux(+uby)T>upM~$jKj2!**$q!(pM(AcPU3mzLg)SM_z!~b z!B3&zdF-DTvD>fS>$v~;9M<;+{%dH5=c}yP8^iVRR{YXIzi)Uy))K!rq4AOTM9*V) zU#O1uJ$ruLsGa`kPknw*cin1@y(@edjwRoNlw&;I7rX69h<*(}$E){F-YXbC)W`oZ z{GX=2@@V7el=xM^?t1KaC{DWl>wK{t?o)44zU!a!!ga^>!FAO0UlPiH4i2Cl^C-{t z^JDDJi(RC9Pv*JSbG`lUxxG92`@_}nXPA`ol0)OWb?7v-&+k3n6S$rqB;EN{jC_@# zabIidvpnz7y?@+}{|V@RR*CYb!}^rt`>Bjx2XDb|F?ts~4)3A-hR}K+Am6(%0{@5L z7w`i4j^g(|G#>9vy6e30ulFJUX$Sd+;=dkkyym@+>(Rsb&4L@C^QZysJAgd_<^6zt z7<>U{Cf(m>81H<7-E;YV%Ik{V{pl3tv?IMQtc#!VyZt!@|LxFyei``&kZ%**3lou^ z8hZbFfb_>H&-Y$>T^H&5`8=VwzmcLG0=16TliUCS^6s+zfJHJcmjR}&y(N#r=i$A z58AI8DYq6hj!jSgBk)er-LD-t_EQPceJ{_D?^o#jZ-Sl*9Y1@~9Vo}R>oEEc+CKun zR@mLooGcK5%&GqbD@pjr=VY@Jp1Dq+WEeY{IA0$ zq&pAaqnr$+FNB_RT>m|PJ8nml&+%fseiXZLNeA>E%AHC5#sz25E%5W4`XK)AlYRyo z7mUT<_vCqdAbt}`zln035AM6hU7zD$0RJKQJHE~LBYvJ6Ucx^f8~~q!t}EvP?fIPa zpJ8|E^SuQ3)k65WE}f>laO}nvkE8dJ?z!6ce2{eGsB!pB!fsqw5xo+=16{`-qrK1C zukanpdnuqjH~HSnQl9a-<0&WhhVUo2fbwpKuAjeCpXWy3gZD6==Vy~|8s&MfawobL zYz}+FyyWYHc0V(|NQT`wI1B#95BHO9T=E^dDt-rP{~XEJ=7 z_Pd|h&%VE3$oD4x&gbmt9i+bu`@?Y3^P`QQU4MR~UdR8__$PwTlh1zZgdR%zFtq*s z3FY;}{uq26K105a=pQBd(xAPUzlhy1v(6lhMpgL;paY8opjHk zuG?ksTLb6g=lgQqw0|8}!SSA!{1MReVmiw6{(1-c4LA#b=auVHLhOs7@u2S|0d~uA z{rLfVM#}ej9iLBP_xX>||HjkCk*=F1DepXfrO|yUr!)MJ^xH_czs8|GPvk?pK7EE> z1>IMBqJ96aAI>M^oA%W2Jycfm<%PvaH_kBrz76}M_&p0f*9=Ab{*1R$VRs&@jmw^( z-1zv_gq|Zk54?<@@kI`_`$AFjm4nWg?P&KSJ;p6WB(W~hQ9a1_<5f2 zJZ${XoAmeK*U)*mmhxVO^+=x#UANvxyAOEpm=wGB6VIccgU4W9=)U$K+VOBJ^_PZ@ zpY7;E_*;+Vo+UjM%nhsIe-n(4-Fx&4*c)I^gm!(l|BYjeOQz%ZE!;+a*UhEa=V3Sg z9*%txtPR~Kj-V&NN8l*R=>eUWuaaH^Iu9I|hp@ZuRkA(M@*ATY;hzlcI$<0!lk_sA z_kylV%h31W=l+okdl<|}dJ9+^yYrzu+V`CaZ9F@Kd#_itMou^*ojLP_b0mqTv5N-e3KUJ|$#cvvPJ~|#;H@o5A0YBG~yU?GI zJ_tH*jT_S7=Q`a4{SE$`q5F7S^hx*!JVL(vVSV@p>Ashh*iU2kJ@&+Y6#Fc+aiH(T zcy}6p##yc_^~kpyegK_c+sGGuZ@!m!`2Xj8gummjE$N-GKM7A_cYc40{a5HbnS-7P z?C@(X73eF`x33@r& zI42+WrqJ{8K=cyW4EntL(I=qmo$qZccE^qDi|bc+(uYFVBjba$w6inmJJBi7t>KH< zy`OU(vHVT=IgVU!KO`UDg7&P2&Z||V?}Bf`k6<8g$SO8wcFFV?G z#r<^$cH@Ct(4*l}_!H&qg{~{cKhCc@_z!`O^M2@=@H@(}-L`WD=_8@z;2E^*lk>oM z#`VH=uMz$cun+7HM?m)#*WtUcJD)Ar=juSZ>#qHu4!h%FE&5}4Gb{m}&vVf~(;w&1 z%i!%r*ADkaku{VbW;d}5y==+?H_PsAhJO2)&eICc-%h-Jm$I*K1Ghi?1`0kB9 z5ABzeXt^MeZa=++eG#<1ucLi$JJ3G&G4xS*2J$nLIK=h9|8Ea)JrOUCm%Az04=lU! zoK|tX+)r`+@Xz`o5wkLmmpdY^AFfV6+{8GGRHUiu+1?WF9dSM#cbEEKf+8f8q z9TBP@($Y^?n&~L_zm@5E{y1Ll2>am<`sGIFojk{X3p4P%alG6S_CrSc<)88{%DGXP ziT3_eW~Q9$g;{9NwPsf8yc3VjNPX7}3(=l`i-oEGTC)iCUT-W)yRJ5i zQO}iPamu^VSb}z6Y2HJ*QL!X`SDU4%C%U{B{~$|a53&sQtHt{$|4Oqg<=$v4N4sOo z^5nnPtU$dtDl5|d8;zA{cT8ECe6eK}^53YeO8c)Ct5JSbtd3tyS%Z96nl&l+pRyL^ z#Fn+m|4&(ma{ejnQqHwvJ?e`t@5ldo;{&wo-(r30zh2pZ_66Ax`;}rN%8Mx*lP?sT zkp55ElyYLqX5_n4Y)*OCiY=(`TCpYd{kO0c&wHiVn(|~D^o_!{wD;d)JL-=v+v6XK zVWh{D9mp3IJK`5pb|&A|Vi(Gfie2%GExVCFDt5;&y6k~}RD2M>*s>@2W6NITk1czX z|G$NOc;2hczSI+BIQHnWAO2CXKYp=g1o@-O0r+2Q4y0Z=2z{kEnDV0HL-<`S4x#+0 zI26Cwav1qz%i-jYEk}?)6i1SNtvHJMf*g(gYI6+rM8&cA#g^m99~H;r7gJ6kUnou_ zJ-VEPe{?w+|Cn+L`9kqw(yuh9Qf_QHjr_6Ybn?%*LVSdBV#}H255-xe|680*{Xss8 z{YrBV<%Z&1(yuk=QEznl82;A_=hL1bpTHhdK1seH7hsPm7n1Lv@+r!>)?7rr(dA6$)sz$DbJ(Ne8vKG> zi#@hnNB&S;Pr7^_EjOV5DPN$RP~1rRmEtDK3vx5|=yD7GvE^3sUny>*yx4L(`D4o+ z@(4R)Z&q$Y_ zql5ecdsMuD-| z760h+H~i)A=urHF^grQW0mk&dSUP_E%LT-g3CJfCqGckqya^qOiAldXmP|svP)tgC zY?+Mw(d8}p-x^IO$4{m}hhj?7gG_}zx=f9KD5fDj$lI|8nHIZDhnDHlQSlD^g1i&E z%zzFuBlaNg!Y(tRgUpOQ6tj>XQ)VS!klC;YnH{^#fsQV7;xF$;Un}OKz94gB&l6x? zbiQaZKYj%QEQk)W5cb%zF!^N>bdW`{7lXwEEP)QilBCO0Xn8L>Dwf7imOez9d2^2@I1=&~FBQL#IIJ>Y`@_C(9x=uqrKdQ=R@FUWq_qs#vI z2N{7~4nPltg903kmP61%4#h5qp<~M7;79YNWQKUE3DBwY$~0O2VkzR&)j=Ik^3I%m%wto1+5=jZ>t?|Hv>pYL~SGqYad^T?wO9D};> z1!Unk=z)!<)J;sF_kTpL;-Ner^ zuDKC);TOo7U!wk1gxrjInqQ;-je%QG7nUM#RosTGxgB-k4&|6~c$3QRY!enI46x97-s)5~5zg4jZ@@=4;qaKd8W4iVpzJYP!o5(XFKnmLikp#zTabmN z$loe%MHX&D-mbUL0ip`PbJa%DA)P=2(g>s&ITh!$o^c#?G0uw>Cf%Jy)^E#CIE$4p_>?@ESSI@HI?dMh(Vlu4 zq;~CxeqG$>$$nqle)Dd%y94|@)$Zq8Q@%R!OjuL)RonN$eqWd;EPWMTn7eS*OVt5+ zuEI;T58gMr4#;y8Mz;#z#SX9o>;OB!4#c(t^4}$4Y**)M>;OB!4zL6406V}Aumcyn z1M+>H@M7Pe_hSdx0d{~LU?u+v0-{t!RXCJ5b zj$IIMZ`rxb`)Ffzu?7ZRg}JUgkK65bWq4dUZdaBk>8b33!rc74v|gR8{t>@?3C^8M zOLF?vweRX&<~t`^9V)+bm#kn>FsU%RAU8NTcUZ8<-bH3^R$8wc!*9fs(k;cGlH&FI zeZJ&me`;!K*RF1_&!6h``aP+sUXM2=Ioad$IJ+F#sDDOBdxy>yNU!5;){1wUWxqCU zm>g&kXIaj5eR(P5*Gl{4Y`V+ij}6vKtvSADl=E1et9{5KmX%{KUg=U@2IUtF2^MAI z5$F|rKK;aDY40r+RxCBHSXh)_5X{NWv^$!UHzYf+NH03N`+=4X?8Tg`ZG}bq)-!B# z;LxJ%{DQ3P0=?F}4rPApg#88ESL+IEeP>$h+#`4Eg+l)-iHZvaPSvTrcjp$epKrzI z8Ty%23l-Hp&q7|kQ0RwBEwti%3;Fayp`Vd96?#-=TF9V7sK<#FR_eRbu%*gBVS)4p y4LW0iwsEfimrB6zg!OI9oY!f4dFN%nuc5=UYs#}@dLX^NW0I9%|KrRrpZgz+30XA& literal 0 HcmV?d00001 diff --git a/input/exp1_scalability/vsN/farmers-si_n82_T16.pt b/input/exp1_scalability/vsN/farmers-si_n82_T16.pt new file mode 100644 index 0000000000000000000000000000000000000000..e8793b30886f0ad83d939223198f48e5ec5e720c GIT binary patch literal 21899 zcmeI4-%nfT8OIGIKpY53nuZY4)PbZi3BPQ9cj-!4O4qqGz|slj#xmf7HS^;&hiW*o`+_WAXF zp7(iw__)|-Hq~${5@~LZyl{0!nj#-%3TxTIR^L`Gy*{2yKT3=aEN8q-|K^YL(^CzR z$B!TXyne~sC@ig}SF#&xS+9^=>MwfOf1(fUKPJ-i3sd*R)W;Xj`AO2*<&|tY<9UVj zQhv)cR0)5__rC8tN6o2)^ETYeu5WD={!4JW3c;hIe~F@Tq0v%o`4wkPQ{_Nw%&cUi z@8&laGx;BiWSad!oOxoR@dewe^jcU@_$b77&u$ED3}b8+T3 z4W4;rE)_}D&wN?;%v|#SP5JohXG=>P`FwWC%WbT0MIRQHvxV&PgWQrAHJ5wlO}l5V z%$trP;V0(ms-J(gZMBpjVLDeY&J<32=9(p*PhEIsUQ5-b=AW7CpRL;H*F6)PH#dsW zk4)FD=%ndhy_l-=%uPSIr>7Kf%QJ7xo7)TZe%;=CCcZ0T$n^LT^Pag=O!w5>P1QX$ zy`{)L&-BlmfrT@E8W_&B1VU80dN1r{N482zT5>|iwsvPVDS?#yI~j1^wePBymTm^ z^Spf4*6{SLzjWB`Wnb zz|$gwcC_}YhIX^ts~U0;gLWU-%N{ZIs)lDvZV|9WWEZxIU86m!?bTL-7to%?_Bz46 z`*5LTjvN#(3eSG=6^REgiChg|6~D_0+eLQa6|o-_JA@}JUX{52HSCnSG+q<^Vc~1S ze_inPePT@F!r~2yb77a*HFk^sW(!h5cd&2SnC5D0(mHp{e%0_D@tbzwyP~fZ zeqZ=Bz9)JY{y^+5Oo=@xzAyeqg)@>bD1IRRE_^6<7tV@3D9(vL$o}HoIWKn3PwtPo z&pfB>_`}2d4>IeCc%277;+QASTk60MJ$`&%VEnC8AA0h^Pds%;#h>-W`z`snkK{fd zJ^alQ4?p@k(Nh=s=!2MdKA-b>n|j*Bp;wUmf=Q7X7yad&;69SPjEnm~&R_be@uzR-86Wj%UiwR2G2v+v zWWHKN*6W2loTq$#!0TjHq?KjPp|iXZiI-f+Lf z{1MMQ@V-v}sGq*UN8P6+59<*g`pfzwjyh%~j&m&`a$JzlYkY3#6FvL$j>y!@I$=H- zALG#Dq(1t^`saOxIQmar%)_MQpAe)!dVci(Z(yC!SJnmn(RH%U$#Y-wG2iS9;$mXg ze%Q%FJTi9XZ&G*|KjY@SAf9uU_aEj5KKh71>!Drv`8-Y?%roZ-_a*cfK6ub06Ng^& zQWxW-e(GcX@x#uz;ANg!hvZ?tnOFFThY$WYr4R7xaj=fjlZSH_9_nGgu%Fl;_%lxG zfEOOlrL$5OdOqLMUwEjeUi=#cqk^r9hrVAHJ>!ds%zEzzI3==oruw1}NNtaI)ISWkLIKUc;Jnxy10+P&iu3Q zdH>bt8~tFN&qr8cRPzPNa?!)O0;+V${nzyFipc>CJx zFQ*gEznmtfr$)m5<+St{Tq*kmYSe+c8u6UaoOGZT9XMH!ySsQ){Ro@qu>E&;;rOJm z~h|p0Ifi+h6zLxa6?otaY8W&VNw8 zgT@E#PuM(%?H{!6!}be1&i&To%Hz`4I!{o(TG!*s8#d2j`@8BqYJJ#o?zcWy9+$q> zd4lrQx*k_vS01fL23>sU@zaQ1BYr{k_z?$cJ+cdR9DXi6_FADUk4sNp>>zpM|PnrPV062LGAFt z16I?!?5=p1ze^8KH6$L?{#uU=;*Shg(_?p`=5hIH{ZYk{$A#2|9a$rG5IZu69z?GZ zyGH!99yzStm3U-Nxz_>l2OPaHCNv_CT6d)IN;;o)~9TJMU}e#rRiexTQWcFxK#>LDIvob(Ys z{NceLWS!ureDAu6p$F57#$lOO^%HX4Gj;EjgOCyj0`3x#>SG#iNxqgVt6PqG(IspXsi6< z-(UYO?pJ7Ev#B~;?9P)T%OAqmX$`xsJ)R?}iOQd)*R_k}D?eccBJ~76%~$?tJbuhk zqMh@9(^QG?h3D~U%JL`O@q3vP<)154?mS^nGnN1HgI^Ann119h{0haDo|C6(%HJ?{ z(>%Z7rtu$LxO)!E&;Q*tU(C2^lA5M``|PH<@`0OXNYj+>34Tpe>P>8}B8}`k%~QT= z`SnzZ=hu~aD(|P+RAXarDK2uop-=wn#Ru~JtGF#9k@5MGiyv1!dwM=B$@V~+O`R_O OE-LCj&d>UneE$Qq_$o^P literal 0 HcmV?d00001 diff --git a/input/exp1_scalability/vsT/farmers-si_n82_T4.pt b/input/exp1_scalability/vsT/farmers-si_n82_T4.pt new file mode 100644 index 0000000000000000000000000000000000000000..5151d6b4ba64329a3f1a578f24de0bf372d50493 GIT binary patch literal 14017 zcmeI3TXS1i6~~Y4#C8-rPSP|^;xv^VCy^6hbn!K`)DfvlQA1=C7!5C~EUT74mb~TT z(lV0)29lY+b%udqn4zzI-~%wc@d22D_L&DR55N--d;#{B*3Xr;bu1~7V+MR?j{mdw zUhBUud+mKL@<}FoPdLuNfV1zo;tV+NWhz^_%FgsoAzfOEq#w=AZe-ld%=UMS@z@FH z+xqxnPu4A0vibC8uDq3VD~0S#RmA-9-#34hGx;Osd#rkbsWaM0I=8WzOJ`iSlFk-) zRBsFQH?{OVE$mk(QfH02n=9>sLpA1^^sZYd?)cM{ z-1=^zxPj~MNBrtcFqrDsT4GHN*HrxPluMg?HqWMfS|p^-rFu1Hw|CX~2fyfb)k~`h zC(-lZ_r9mW!*x#wBvtR*~nFL8}|!Y*RL*4uBj1MU0PF@tArn_ zD|u}{Kb)@_2&>Wj`3JxBx$3GRUQ3*Os$NO>5^GP@wU6>f`&Cy3*3?*4`;i)N(2l6< z`SS^_JfW2*Cu<5fTy=9z-AeUnzumW0ut6iLCbh4Ce^!8xFnQrERlt4ODENqsx zba3(5WcyAm%ozvX_@4Tw$+bDX9T;`yn>eukp_?n4dWpxv?b{_g#e(Y=bH;}!umA0` zSzhM!_TiE<*Te_2Tg-;yu~5ejXk!N+W=hOqqsf#xywI#TbLh>ndVnXT4Vq}p(F)CC zH%BXEBL>YnFo)S<%+U%@*TfufP}*KNB;z)kt=b&JHQ0w{FE+<99-N1BwPELZ@x1Un zE50Q9;00;7!k0zvqJ<;U_QFdt{=9fucsj)^qW9m1qvDs1SEc{3@D<^|Civ1c?}alm z4$ey3#*p-bVQJeKk^ZO$&Po6C;=J&H3({^E7e)Sm1(zfqFI<*!8}CYgxA05Cci8xv z*y$F&E_|)<%OVH9A#EF1r2kp*p6IuWF_Hg&#ji;Gt?<6cwZb<=F7APENq@KStHNjF z+tTlaUz2e!Ovre<_>RcGD10FH+QqMnyca%{aW72Dc)PeN@*wMrednxvHlJ^0bbZ@Ph4-Y$9*LC`Sim- zAo}pr?~{J~!X9~$;l1Z`p11L5SOlg7xi46eHgS<(_6hDIu}fUs2eSW?SL$F~c)6Yu zFZVT9Wt{wQUE#T9M8*SxA;EKkcLd23*BkZ$_9gN=Ci3Kse&WL)o0t6JS3r3B1(~lw zY1{RJ9rjb6AGx2T9++q3=pPkc_C4x=`5}JdVjh`~TcVGD$e{;+MC9<7{f7G`<_~@5 zf$KW?!+-JyAAX+@JJcgQVo{(ep2VyxhM9RZ`K8R0U5XD7{?C! zv>9jqR)mN6iJScbefC+dKg)*nnSF)(67mZlJoM8>kA9mMzlan6 z@sIgOj&b6GmwBcRvBP{bukfJ{AN+SE5AfP?P{;IRhkX_v{9(PYo>(8q6DNMa3lIC! zY4MAGo^Q!7JowWi@_mAS!6A!>ykC@l;tNQddLNfI?=e`P?C1DFT-4u^=&`S?NSl3! z>k58zomv&Sq#%CTaoYK!9s{BW*!dujBO-v_gtS9~D}tj*n`yyyqS z59$m({9!*8vAz4RUg@ysgLN_}VbJA4%R8W9gaf^p`bbCfZ!u(TT>@%z<))DIu9@ZcEfrmWV^@2S7Q<4Yl5C?oKGLC)biFJz| z=94-jf7Azj)CGRpa`eNCzcvr!>`TMqCwcMWWgI^9%jx*Te|-8oop^H|^vmg8uU}5X z@z_Ggzns>7!4qkfQ@$k|Hr(q z>J)pVaN5;7Z25NmIBYrA(_!23;_o!iUOeQ(%Mwzr?h`T)YFtN*ZIqJeHzzY%Uix9@8P zo3Z{gfzbxW&zmyteV$+w)xS04tzV7rTe|4?yH@+gz-FnxRy4AFtkdVgqCfg*oW1&; zzmet7`Ug)BmWa(#zit{?Qu_6%2TRmusb34cjjrY8UHzl@eim7AHcfr8nu7h}_`E4i z^JSDw^z}`N>a|{8;_v^A;qTSO;5Z-1`)sXOeY2XD_xxrXN+wQL-zNL@M-WE7Ci@?? C8z&P0 literal 0 HcmV?d00001 diff --git a/input/exp1_scalability/vsT/farmers-si_n82_T6.pt b/input/exp1_scalability/vsT/farmers-si_n82_T6.pt new file mode 100644 index 0000000000000000000000000000000000000000..82cbbec4cdb32d1dd93798327c2e29344b1c481e GIT binary patch literal 15297 zcmeI3O;B6c6~~{z1}uMIJC3o9;{e97ut5lch3qDdOdZlxp_ZxJ(R8Bt){x@kL^Y1h41IzJ*Vdguux=(N0<$De!e zIsfx<&%N*Iy?Lo*_c6!m>vMJ-SDZfQoor<-U)h-0C}c|UXy!?5ayjc}!|OjP&d(il zexQ$^cjer2CAXSc$(PsiZl#b5S4GU9;6w9AIb*+2fv2h~ls==4Wb(@^`ApVzE16tz zLv=S%e?v<@(857=EPd9fyZO>ax$-r^@g@XMs`538o^+3)*w7y5Rc~WOGnie;2j47~ zm$Jq0N=xU51coY`~>#f@O5 zl3&^^6qj-R-Dpsq35C+VT1%{{{+deg?Q&^l%jVg1K#N4wxpcS2-1??EzwqmBSH18! z=_I=r{up?!F6h5WjW2$X%ax18e9kSDOB=yQmF0XTzx=R}bA#&Q*rFP6)ulysxk~uH zy0WV6ulBFj3{0uP)$&82WE+1vNium1AY&$K0Traa4e z*oihM=kPn_%39XVysN8RXKsEjT&-$T+5V0xTf33=Ja-8(7DsVfu;A|eA(2?{M?j(yJV+WaNT0w`0(tt zzg{-W%beamTyka__+WO6$;kX%q-6)Rv0V={CFZcvWXc>~XjYs#^yXMSz~j;eO|<4{ zhGwywqZzUhgJvC=!)!6;Xoe?iVh(sp+Fp2C#%(lPwK@80umjCrY>p$`JrCzF*T2B>dL|U*01Q ziC(KXEP7rzBI7n*m;OEAEj5gsrmZNE4x^8YIsmw3GJu8iAwPx{-1-w?h| ztHN*QN2fi+C8xzvMS9~D)esNCZ|6lQ&5`Qy%D00p44UwDoz&EA8UHC2G zv+*tI_rh<>xECg6+%LW@@&|G*<`<&m~_|q=}#G$b8%sef&cXJ@}&{hrjGM+%GYI z=ra#o*U2CLlQ;PA`>k4}0AHSG~gxJ3;NPg`3v7a|kC*+m7AV0RB)H!w@h&|?;b%EZHjN5XI zV+VcOj5B`;;URwFX1_q6eU|GF^8+7wM4oyW5Pp7-;|KH1zQTP8`GpT2`e~y_zs-wZ z#EJj-$NVG5IB~(tJX43*VZNDH_|S(B{(F)KcBhyZ${(vAow1X;i2vse1bEA>Kt#$|j` zka)SSFyF+@Iwo(T=5U{JLOPoHoe(+pU-F3_`FmC5@C!d|Ire4ZW1f+-^9Dcb2s?JX z=nsh>)ERpC!+t7aTlZhx(qYdB>*SP-V~0BDK7e|%<71x~5&i3etY`L@3F&7YQuoA9 z-Luay4++tup2$DzBPHY11M^56tW$W&EBhMpT@W7PBQEw6;(>R_@{9Wj#+iTCJ=b4* zelQoACv`^ts1Nw43;eX@=!X}7Z63zim-@v|^5VtIIDF>K>ByrmfA_IYyg7G!b9&F~ z&FR$qT&(3cr?oe@;`R}=DFcBv@f^_{WuP4yII73(Sv;tGv}&i*^4qg;c&%3BY`3R} zZQqXPp!{joPN(JVJRDZL)i~Sj=V9Bo<2fjQTD8+@c{>k>)owM;cKg|G`}^hFufE^> zv}&i*@_zm9v|Ouk?$sYJJ6`?mw&Rzt-TrvlZPiYv<-Pno=>AsY+^c_HcD(xAZO1QP ztNxhhO~3l+lKW5nfYzM5J#PwnJ#UK6&-KXjCZ6vZ-oJmg@Y_D^xH)&@ZS)~V^$URM z*1x~t`CF^lB?_rJp* z>Ojoty$&urV-42V)`v(p{5{Z}4jPMp{>{qwOy#UMx|YrGDLPXGwdrOxrB=Yk{|MwY>aLzjWCdk*qkIroLEB z!G3*w+K{I4GD;B%j@GO0I2n0wWwxPI@_6-X8`K|x J8Tp#*e*od~Ag%xa literal 0 HcmV?d00001 diff --git a/input/exp1_scalability/vsT/farmers-si_n82_T8.pt b/input/exp1_scalability/vsT/farmers-si_n82_T8.pt new file mode 100644 index 0000000000000000000000000000000000000000..85cda61f58d6ab1828a3e0a519cae81a9a0bc98d GIT binary patch literal 16641 zcmeI3OK@9P8ON{fBzBaC6E}^MI89~8N#w*YNq&ZwDkgO)YKUS2qhX_wZPnHwOFr^- zq0D4>O)}G6m!%6}$BqSymKioIm|+9MW8D=jfE8PooFjcdS3cIUu58(6z;{Oe&-u>x zKHvH7m9I`J*>S@2y1Kj_*EO%pdpBEN%au1qHu9O`Of2(sW^6g@XGhn+RhXYU;eAtI z-)~>?OXa22%u24bmh;Q`rO^tD`4f6*{wQzg2deFvYLBGPYa^N5@=7j~_5E^Ysj#6s z8i>ECxgTiekUEjRV8s1gaidiJT*1i(3Z7Q@t13>VPZ<>(+T)z+tS@LLvn#pKTZPhM zw(vC>sV*JF>1QewNt;w<*0SsC`QnN?^GuzMq(hoi2`gV*&V5Int7;z2Z2I}aMkrIx zEpFxu%eektETqmyBI!;oCDv4TRV4IIskpLb^Fq2!vy7{Y=?;xc>zk_Q(MKJ=dg)2h zOSV7yecKClN&ii1eEEZ=rBb1gTk`Xz;zsCkc{x|kEkDdJ`5|?AXhHS*>dJz;T2c7E zy0)tAuXe9i4MbJnYR{v;w)^V3QM`~m{X)HxY)dY@P&YnUHPWy8D!icjE7DKZz&7ca zy1CkuZ1dHi79JX^3f%J5?FDrw-LC!i-&T=r5^*)8B^G>jw_@(Ox|eKwu7<0UBfc74 zP-E%STC%vgmMNB&a~orQFp0;E@DQ7wN zd$Bs@Jo;{_yq5Jd@98SnnVX**tyHzC>|j^ac<|}JKKqebD)VuBFzAic^FT@p-B52jm;jnB`GH|>Bnw(Dh<#9TI-Oqt7tX2+RJH^<5ao|G|YqBU0| zG@IRAjgUqRntfm{(_+lk2+vg6Jn*cHU3gCBZ8WXgT-{aJfulk;>!^P^fw_ofL zJ$uEML>{~(<3{+h@LjgBSH>>9BJ=yjtD>h_ye4x0ZP+J%*?3*(2ZgVQ{u_d??h(Tx z*DUsnoC^nJ-o~3k9}I}Mgs&0a7Csl=k@=t)5q=wogl^+qq1$**=#6k#_-q^zx(i2T z9vqXgjpITGqcXNJCiJ)iCxpIVoD@Cal#GMowDAA0;EcrM!daQO@xIVog@WJx1(~OR@_fv5=6P!# ze)MqvVN5+CZ`*+nIo64Oiy!bohmY?I#D7lwgN{A)Bah!9;isOs-(rvFNS^bdqrXe! z(GR^%==g;_@*vY&-_QBJjX&MOFf7P(K|;pFMSkfMJV#=exOfhv|B_egU_kV8KP6tC zYp%;Y`Qg68_m*Cn4-1Y9UKG41NS?Ug&z0S(m zt{3dkPx=1H^Cb1aI>QIOPxR9Fr~}rA_=$^kWIgVPJpRFl9QtFzhrje2o|jlZ*Npr$s79c`-Iq`9??U7sXyfKBPDY5wWy3Ef_z`&d&7v(*`If1jK9p^_PVaJJo z}zJeV2$1m0)A@*kl$&bB0_Wv8G6Y@%3kRRJm>Kr={#2)L-zCbQ4^EMyz*g>8# z^Q>P&^bkLB(=U*x&vO4^eb7f9;in#YML*xi@q=}yukc(#e$j^>=!}trZtKM_;>3Ua zWBuV{p19DF-ocMuW^w5{i zh+ok8eoKDQgFo%Ue@ZYUc+S#8-Y*NC_`))#-UnpN|1sE~^mF_mF6wVaJ0Ik+R~V2Rl<~M=LXiDSK0Ad@Ua1%IGc5CCg2c;x zh4m(G_Az-AF^}hr(?V$EcS88+zvL4+^7oqX;TL|`eDr1FW1Zo%^M-!*5q9i&kq?U> z)ERR4Lq8R^t>>=}A=vA|J~=D%*rCpO4xpaw_~;V@B7akm{Y-xu5jy*jx+i|>o<74m zBt(vSBLD11)JyN%RmOanVnR2fbm-FP0zYj&=;*~?TMzT}rEc+)ytwo-k3RF}bl~wne*PPsc=Orq&FOvDo73q0 z+*H$VPOEQl#qDEgQ3l#tq~}=8aRyqEf#Z7Ip2fq;N3(Vgntyv1j&9U!oUQissO{VF z9F{-L+Bs-`I}b-SZZ^(V`+3y%?RXB$pJweGG{2pPqZ&6GXRG~ewf+6-+pm1k{4{Il zp!tLPeb9W(#<^F2Ty|XaR@(`xuhssz>^5uXp!r>X9yYz%IQQzG%Z`iQYCA#owb~z- zU6&o3&KPv*gASjK%-aZGP&$0bfi|773vD_0Ty*AJg)TcTI(C@{p}VkIKKNbeva?@0 z_G~>ioiS+h+hh1_wD}mr#~5_c?RlHe9>Z^=&BqwD`Gbyw%DebnbnG(^LU&=aeDJ%_ zWe2+XfBfv)vp-$v)n{AuXSe^y-*x>zJ~lshJK+E0yTq`u``K@Q_lX&s&u(_>9b~JV zpo%wtevS8>&0@RwZOidK1Uh3EKYXCg&zN}|UGpwF^2~!aov{mTIo46q;<@A+(b3b0 z&wfu0f1^C(H;9}Ic`xajhmJlQ*@w{KccIw_zy8CIfB2ye#C+WQ;Ho!NXa8*-KEe6jtT}+2$uG~(4F+-`IyE~rJv9}Z zo{q;SCZ=a*W+o@2Q}OZm%=GNU#NQUADME|0oj=CMEVQr1I$4vXS zIrx{l8n=$O+KOu5y7Jbzs_z@R=y$qSJ)>%Csl8Wh*YcrGpF@jY&F!mK+xfR^3F{xe zIkd!VEw%e*yOxjKTH>~r+P%Q98LD~tnf|fBPG4ll*=lNw)fDXKSC4)Nzf(>9ZInu$ yI>k=}ZglV?ke$b1|6bV)p7$$xzpv#g9p|Tss$J8DQpuB*Un7L{AzH&<<^2x}iYd$h literal 0 HcmV?d00001 diff --git a/input/exp2_timespan/T3/synthetic/ba-sir.pt b/input/exp2_timespan/T3/synthetic/ba-sir.pt new file mode 100644 index 0000000000000000000000000000000000000000..8811e00128b9a28bc96e0bb6c76a0f9d714f1239 GIT binary patch literal 178126 zcmeF42fWYa-|&yJva?n8&L&YQn`C6uv{&Zgkcw<7qG)J{gfu8gS}Kb6(k@D*z4z|n z@BKc{bN~L&b$@@qbB+_G*XzFD_jP^7^||)={*Jq2+rzUaCFRVS^k4twPb!quw@-~R zLq^x?-zTL{%~3;#wQ80>>CFH6-{iykrHmZiZ*Z?c14a%XkTQBmzncHi-uu6Dwb=W= zl9Q@VNzO7c`LHTI^X)05*MR^-yZbung7!~>FZ{QnL23i}KikgN5uk$w9N>q06y=blO& zF)=w;m7aT>s@L#7qecxGF(^6r#N<3xdgj`b^`B}D8PR{h_~g9%JlE_sHf6}LF}Zq; z9?*B}kYW7|eO%pK$@!{O>6v4X6{RKT-)EAm-N+GxQp+sRGs_;8+Q~=uJbVxPjT)O= zuzizLQj!aGYMa#du=Zt2&rB}7=YO}&GBdfzrTzMi95!q~zmy>(M~umJ`sn@xMi1!U zVo1M~T**bNc1|vql3cuVa*6*qyde3g!F$RdoPY4X0_r3mJ-A@|CHYg5kJ;n6{l|y5 z%{nu=WZNuz{;!$Ir7j)3*S>U0a+%J_W&dN}HM!jXu&@lxeb)P}S zl;ldClPmW;Y)|AVZIY|}4~u%qRri$m@!=`S$L=d}Msl^bS!N_x-)CDRCAntjURz zcW2xC-yt4qhrKPZ=P>%8{|@m`JM3+Nf9p8Nwy*n>4sl;Y9BPLHvmN%H*7yGRZ{^-g z{NDc#@lZSb|Hgq2YlpaReC)4w*!$kI_rF6t)DDNn!QVFz_CAI0{qGR>Z#(4P=RWBW z_ccU%+aV9_d5HTO;!ryr810akae0XQ8X~jXVecpFz5g9z#7x(^FVDEp2 zxUU_K+=rw?%-r@UNV^~6zJ^F^I~1ZF4{={Z#M=&qX}5n0i%_3`8;er!jAk+X9_TDi zoc}hKpx&9yqbT=aq5hf8vXuL8VL9rX z(JYVO%w`43J;Z$tk=}NwNP8YEtVFyA3o8?EX0Zz89jL5Id>PGS@taw!MtK><>iGMo zvIg}#P+61sGMlw1H?vrq@(y;^p&psdx|I7*VLj?|u=6}b%;HIucc8El@njTF#@~U;Q;094*%-eE8k-Q;fx@Q5bFi=(@gAsb zPJDkCTTp&xvnAyotZYU68O_%C-QPTwa{ey1q5RC|X_Ol;+oGRYY)5$q3fmJ;MzaHc z_ZLsc|AE4e#PfHv6LIWscBY(+W*7Ya)7X`IrI+3Cv%lG$at;*sAfCUAJt_b1;u)0x zPvM!=C%rriKk@Qx^fQ|0;5Sj8i(N*u7k<;q-q>d}``~wfu`mAj7yIEqz3h*l1C0ZS z>pqOO)fVi=L zh&Qdg6#GPZ8Fo>ejb3JR4&`PvFURlx;#~Y6=$uEKe;2Qy{CIgK`ca&ZUS@FtRz84zX7NGF%P2mCzeM>kc8PKWcIo9K_}O246#wz^G4$i*M)Wg_n&etP*7e&XfR=%V4qQZ7JnJd=kOcF=g~`)FJKqqi^xNK z33(J>MsI&{EB-T!uTWlE`6~8_@-^%d82u;~K`*^5il2B{4E>B|ar~y2CGZpCQOH9) z8hKiI4E9khiCz>-p_eF2W0xq)V3#P%Vi(17=tZ$SdTC_^>=R{0>=I=q?9$50*k=~2 zP+p>}ie0=s7X7rc8usaBb^OH38t8{u6M0%$3;RS_8@njhK`*VWi+zaokjKm8&`&Gt zV;?V%M?YRRKtHW)i2eTJ3HVPhPsC56JPErfHbO6>c`|+zfo1hnBQ{?Go zGyKHM=IF=E7U(C+me?iAR@lYM*67E}Q_)Y9ZLrH|o`&DFvMu&eY=>SH+oKo74(MeR zPsd-P?1){W?1Wuf*%|w^vJ3Vhc10e=Zs^6!?&wFc2YT_cC;Ex<4D90Nndt8?o`wHJ zc{X-YJO{mac`o{CWiRZb*c-ig*$4f2*%$q^vLE*GvOoHXasYN2&4KuhmxIuc;$ZX= zv~m>oX=O6@A&y3#D92!z(M-W_S~(W` z5XT{p;`!)BaXfnQ@&fcjybyVa7a>oS6R?ZoMD*h2#pp+I5_%cM$@q(xQ_xS8Q?dKI zIF0h7I32x2IRm>W&O|SYv(Ss;CFq5CDe`!E8T#qvZ2Tn3IoL(~o!aiOuMn7IIK|hK&qZcoiq8~4BK|fJ0!!E?-$P?uX?9$6y@sn1r#6Di$hJL)f z9sLmRKpx_q$fLLly(q3mFHzowUA(*-{j_oo_VMx_^wY|<*hleR^b+Md?DiM$!+(hD zktfRgv5VpZ=q1Vru}hQ>VHd@R(F<_{@(>?E9>qt|i8ec=;^)Y2|a+C(7rsi6yHKGif^MA;ycLW%w%I(-k@m=(y z_#S$RatC(t@_qCZ?w+>L&`{1E+k`4RddevCZCPmo9PQ}m+v8G7mE z=lBWn3*?FNOY9QmSJ5D!D1C=bUj#H`4pm<_!| znH{?jb0CkGInfXC2;}iH7y22^-1rSK5At}K7yWpd5B+$VAN_b)0R4D*B>M5PAo}sL z5c(k&Mjpi?=tZ$8dQmKfUZO0HUA!!Texf`IyAY2?9>rtOi(*Oi;$XCfdmC;L-Rj`X=*#f^IwnQE;TcIB>TcaPvQ_%~t z4e~^J8g_}YEp{QcLmtKU=tZ#udQm(by?EIX{U~-qFT~EsL+pY)ie1r*VmI{SWq0(W z*aN+I*%SQ`&p;mHnaC66S=c4Yv#|^D9OO|v7rl7d3;hs#BTtllu#1;{(GRg7@_5-F z{U{DVFNy=v3vm$gv~n=^@p1_IQ9KX55QicUaTxL_4o5FujzB+NjzmAiQOH9~MxH1~ zV;99S=!KYqJjAicLmY=Zisz%3D92+LFE2npiWj06FE2tr#0kjbUe zlaYrw1$l^5kw^b_SW?4r0Fy(q3gFN(LK7vf6fA>M{O z#M_a_%RA7I;+^P)xC(g`SECojyU+{qZsZ}ZK_22g$fLLxy?A*q`XR1E9>x36i{g6p zqIf@gAwGaSiVva}#fQ*~mk*;K;s)gL@)7hyd=zrJC z3-MLtQG5-(5MM_gFW*2v#5a+L_!jaI-$owdJIJHB4ZRSzBah;{=!N(m@+j^=FU0qe zhqx1Yh`W%7_yO_|cOwt+L*!BX2)z(LMjpja&1q&5c47rF(2{}^CJ(j0P+xzL>^*6t;Pg*Xg(6o;c1;t1qX z9En~ON1+#DGV&;nMlZxM$U{s)9^zQ!A&x^H#q-e%aXj)6FF+o}3(*VlBIF@XKpx^m zyd|eKk^VCKpx_Q$U}Sxd58}q4{-zX5FbGv#YfQ#@iF8fZbTl%P3VRAIPwsmKpw@- z=!Liid5BLU5AiAFAwG>f#AlF)_$=}epFm0 zaVPQ+cOeh)1LPs@Mjqma$V2=Hd5Bs0tey>KhdE(xm=6|(#bHC(2)2Oj;V5`9yb&&f zx4`9aCA<%=hxfxx@G1BVd>Q)xXS@vmZczt#KD-D%2p@s|`;7j(U)lLxLS9%I)`bmW z4>$nc1^s(x*6;a$HTesE7*g2^@kV0iO_!!&VPTa z5OFtyJ>dpehj<&nPSCg=KRGDhdgLLU7v_hhU^RFG>;OB#ZqRXb3F&3<5x5P01M5-G zCU7Vm1ykS@xBwQRzFpxwxDNU|9DhgS?_~1R9>rl3crNS(N5MHTw*8!cwqq{pb0jPQ zkAcU+>aZrP3+qA0&ncu^L)*It>2u&{cp+R2mqO?3-K1mZ=X2z5g}b2r^CRh_X!laE zG;9P%kLX2iO}9fzzS$c`a#wN9^x;-y**l{n-wl z4KIW9;C#3SJ_g-CX_h3ak%aps)RLTMvE#(~=6G@56~JLZSO^w|MPM;l96HX9Chd4D zO}Z>B4=ckeFm_zlCcgn}44cB!U|ZN3o(a!_{h;IW8q(Ln)$k$s1l$Z=C*C9dHOxbU z7ltKaZFoF95jKW>;4tVu<2byW{43zK@CJAbTn2B4cR>GL=vPU<3;p*3{dfKHGf;}c zTF`y!6w<9=7kDo050l{(cn^FKz6f7~{?6Lp?bc+Vo(lWJq3~LGGh7bugr2YImj4?6 zu9uGg*nKh|hK_T`v-?e1`(DpTRHT_t1qR2jk*MSP;6dRwZ2@Hh?W)C)f{;gyZ2NcqjDVOZVUN_TR&P z1G)e1@q48G_j4VeO&CATVQbh6j)0@!cz6-K7|wx9pySlP=kWK;KO^^dsJR($MPLc& zzU?|bi2MoAcRBAfb_3!h$EWLb4iqXv*S}t*$H5EWbm+KRL;6ElmiC+luYxziTcP{j zvDCXhYz)ta=fiC<3l*shPlLnZbk!Yg3x{2;v4 z_1tmd`t3OIJaQE6s%$|&QJTz`B@*i`)ni9&0upF+rMqdcRt(SoyhM3 z&xB{g{_s3F6eh#5a59_%9VfBt!WGD`f!D+3@D6wvya&2&yh!?G_!@j0ehNQ>pTlpV z=hggld>!b%-jMW3uo-L*yF=dx29iDxUI-_{OW#UI=}sSw`CTW#4&yhxI)b zwbXNf^U!t9`E?}ySP(j&%8-toM>WZ>1KsaDzqTd6J)8<>z?%lz z$#tVQh7;g4Xgkd$y#_uFx58K9HuygL9ag7)=R)5t&L=$q&V;@@T}|3~;B-lEku3Y7aYGO$Uh2to-1wpz#6bN zYyz9Y=CB294ZFdfun%-yPbNJYrofBfe7FSO2OoeN;Oo%y+z!&8K+ko>=%0$vbDQr4 zEy-^S&w`GRL8OPk5pWE2{k@L#V(7cUO49d3?`6IlG^Jy`U!6g^9~=TFLeCxBN&9zs z-j}K{59>nTgFP3vBELNx16@xqCGES|Y|_r}MWkJ42_q?X9nDQY#IC1~JMRgOuM^R4 z44c5FusL+Rd7kV@es9&~yIbwEM~MTu8Y8wZ3|OEQo^hz9i}D&~`kF^dLAC4uj+2 zWH<$S&T@W!K>iQV_YT)x=a=(m676>aeJJ^G2X>r5N!3QUM>UAFUU3WBT-*s(=#pEx8zDpINJv>iLA$<$H8+s4>h_vnIx>taB zJ--$uT?QTt-Pb)g&L)2ybU*Ult``06JocU4`_NO!J-54lxsQ7;%SGJzU{^RgrarOj zx#Pk9^?Y5Ecs%4bXA%8tF>ZvodT9yFkyczK4BBel_~L z3H07Fn0n*>zs`Y<`{F3LUwugWO8on~hROK%_dv1Z!*$B`)Ze-IbG}?a+Hr9k>8GLf zeSq>kkLRO&-!(Cg=D@*~N<6L(`^UL$X(fDr-+rnP39~=WG!gpau;^_u`Z}Ysmp8QAQ6Yvdq0r7jj z_jiSkXU~1Mm-nLbl<$6e3hC2eSLirNA#MCqNnZ^c6Nl@$-*a!LUGIaQ3rME+r*S%e z%iy;ntPDK|97p;jcrt7ZTR_({*Kx=9eB_JZQuqM;71}R1(caI(S78I{eJb?3*nP?U-2Lfu;`svljysGv3Q=A!=>6jw{Cf^Y{9pID-1uz@+rSR+ zLHIoU3m#7SRbe~mdC7NT+t2&$2k4zdK)z2nFISQO81(lBwWv=IcrOf%_Y%Ze4xR{I z_a>0O3qAo~g?Wk7chB;qdqMB*=aar3ZiPGH7w{{XM7_#Ezx&u8oyhM4hry=Q*LOU} zkNcJ5r!eumACx0q8{Q0?6Hj00eIoY!dKBePqx>sk+TDG~dF=aO1S_91f>m6QJJ>ULakSdd`E^%lTQF^2$K# z={R#=9g4q?;HU8Tn0B=PocHc8HSu2uc7e9jH0rYedheT0Ii92EQr;_p`dDsR{CkdQ zM7k3k3hkHkNqbIqeO*bu^XnDTufgx(pD+jU=YvN=zc(C9+I`dC`%NeRQs{c~IBCz9 z6={!d@M?G+yczmV_9N*dXs-g$cLv|d7LcEvdK>|rm(Ifh{iCf-{}uYlf{-X-mNnM|CEU~HW3zs}p9#J2@L3*UmD!9SqqzdMQZ9q9Xn_Z`=> zam05I{090SX**vUA3Q{{IkZ_c`AcT^IcQQAg^34!i^!m+=&#eA{IlY3EN*`onggMESPw zQ}_j`=QGd2?kn@~`xyKIzwYDS+fmAGakCWd`tRQ_$M4ooYUZa(D;0xIF9&ML6G{+YJBBs zKlhW?q@RR4U?Ix)z0-Gje~<5eJ%@Ongw6-oHOG zNx4N}*O)jQ-?m#>+9US&(0$3j1G+ys{*NMXf1hw$Onsc6*2i&JpEyr|jt|$JUgUee zT1wh+@;K?Y;g|3iIEVUP3qOZBsfY1@NBU^uJ_(M4#vePMT#sFcU&6oZo8!m%UY5A( zLDy5~m-p+H`0*Z_oj6Lv6zI5plC%k%u=Ir$XIb==M(y%^pMU2l#ej?6g(i0AJjcu?y#RXuSxNe7=g`T(l{l;eEc^hJu`n}n6zW2)+D7<4i`1KtlJM}#Se~wq{ zRU13|WhUu6;7jmL==*V9%4rR+hu+s`Q0`0c75Fu*Nj#0=tq>%+|Cgm)_p{4MTi>H8 zcPf6Lhk0>y9d!OWj;wzz{B?qTU^2WDW~Cg7jdBskvG7&+1$6$pkLRHr znxAr9B&!KAN*uRz=J2hx||@J6@; z{sprVS0{J@bRRjNxMo4uz1a5je9?w-Ux2U1#Pu=h@8D0+`8R;L+}G}*+y^1*NwMWS z{s*8i7P@W@#IL_=yPCA?XbtM!3_ePEuS3VR^{r$dz;)2~M9)(tiFX0?Tvdv={eD+~a*D!Y(DTY}(%#SAuN^Oza|-2rgny9QA4lQG z^}H=<=lLz9KZU+iTt)r7S2!=$Qb26`yRPKZ4;%?EgG=F7=(}D~%JZCcJ891+XHu@e z3$)*J7U2pZixHTd%xqmA%gb!y&U@P^&5UY|K-M?-(wM{o}=vlrue-WE`*2S|6Vu?2fpuC zz`x(+vSC*T9uH4~1EJ%~bIn`i*QXrsF;9~I5PE+oM;!LI>qL9vxe_|R^5ee&YzEJU z{oqZ|_H|xwpnTtRx>5chI11hZH%a0knYtfVpq$3=O}LP9mO|Idt)y4s=mGdTbYF9x zd494VZ2xYQ^9Xzf`aP}#GxsV-QQbpK<>bny{Q`U*zbKw``y&`u7v;7pyT{}(yzpn zZ+m+l_PlTzeJ>f-gCUpI9zB@0DBd#amm#_!%^n>1;^^*@h*P$MyeZM@5 zaxKsHbp7(4kO%*M*ZYRF`@vA^>Hg<<#eM4dJd*NHh4xE*>|IxjA-7+B-{?YqKiD67 z9=VnDoy2!H{2pdU?s>*}dj|QEFVRUYbR80o2S5xg5d1wVwYcg|z~Zs{rF*a`g| ztLt1_;+g;};;#jCyqrUN2y|U}n)J)?3pkqcyeD2vyWS1gK;MHZQoiek-={7n|1#*g z?{v!XeWWkx)$j%Adi*_U_x)PLHx>H+zLB)QJ6KMf{+_80akPTY6W1Tdo%@3A<^A?L z;#dYZ!b+6u{&W)QrqFx-X40-p#}MC1(DRA&!}+~|c>P}N`sBJ>gE*SO;qV^#82klJ zB#tbU+Y??0FN62Phv570cX&E+`8$Uxq_2T~r@5PSAp8@E>pGT~a@Ozbm)ESc;XoX7r-arImGAh zO6HRGJwA=gdg||s>?h~lRmA%w zbbp(T{c`9zJDE7wL*M`X?tLc(`tIvEbH3Y;uIJZN?(NWX+$WU#HS~M-`ILJt98P(j z*S$wLp4}hqzndv%9rV6lhI0H)xr+1_==+d=pXvR_@oPUg51m)8Pu@Gv!(U6{zXCo3 z2`lwJRfmiY(0h>WwwC;d;Rg6D^xfxZ%B=_6!oKiAcqLo`eTT6h{BAXtczv(TLO=LU z_8RFD#8DQGf)n9Ga0{$Vxf9^WFel~tyLx}0u!C}3hn&AB;qQF77`_FYi?hjf2vKYvcTBL4OJK56Gc5%g-q zrf?yA5IzN4Veh-7_op7jHxMp^f557=zw5c{OAq|o-!nEFpejB>ld$#4;Lf9XKIN5hE_wbb9Q*snQ==M8AT zy+*ltF>HdoDI5;_G3Z(HiUjBd>y~v!cw%$82B@s8t{7x{<=fQf#)#qFTMwD!rwM{ zEPg#_5k_i%7QnCTvg@Mb?i1v$!$Wbf7Cr#IH*O<+GUX3|1&G7%Mg9)Q?^wQv`rXud zaV+JZ2OU52NxNS>LfZYiC2@G)bKl-cepTXn0ahWtTcQ2x{$aVZC}#)U30+^Ckv87h zq_3x3&p*eI@A=1bq~|Z^|H(1+*iQNncq#FFUKl|9*Fx7p=Yw$#r~FlLJdCp(TQ^0?>EKN%l*;%jKk0GFhBO6 z!oI|r6}kP~j&yJ6etryg&eLj?@BA-|zU}Kh;TH1$fKO2FPV8N`TsJ(=O+$YM<(`I} z_rg=LcfGJ4-h(S+KL~o?^?dGoj{BwcZ;Adz&~^4%(%vr%5zl6rA9-!q1h$5^!Q0`> z(EX}2emp;Pr+%m<**+;`OoyL9&ppm>?T+OCaGv!i{S+)rxo5-aFdOCPLeF-!pF5K8 zy89OWbPe{7_Y?7Bb===4V?P)E2y5cUb5AnqJovpG?uK8&UtlK+^4@hN>8GLN)q1%O zd5-;yI7(6ODezS2{lN9q_5LXQJD$tn$NIX@x53X$;(ZqS?unZu`>P&)PKEBL$Kq!k zd<>SwUu_)q)gkuPpyTpd(yzi%_}Pe``>`jP z2cIYX2Fy-9JfC_#tVa1a!k6LUlv50Q;~7c%26#982xg(2eDDtFdyDte%J?sio$H?W z68&3`&+zjN^micj@ngSvUzmko_p^_vS7Y+qz%J1K@gDX%`QCq0DEB4k{#y^f-aBIZ zQU43@_Yy3OpQ`Xq{H%q}&)(E;2%G^ghkkFoowWOraoD~CDR(V=7*?bF5zu$^*nNB|daL12@Nx9VQ;zGN>qkBG zPJq*)!A0Kw@LdBIDzsOK=%XO zrM{oLZWTb_dQl61v*6{>@1ox4-X;GC%5k4=LfY?r{to(H2N2~Ig2`|Jd<<^Dzu(1t zXZ1V5g)#n2I^SH!KgHi<{P}(A0t`F{F%bOt?l+FO9gjKCi`}1+uydX9 zyYxNe`+gnk$Np=HU)S%GusaL7UV8sF{z}+6KOQ0N`tAFj_q@^g^BwsU(pOQ>H{dpS z1nuYVdps97&uUWMnb7;b<75Z<-us(U-mmyaoVx!`BV!T#2->giyPh8#ARhpKgRZOI z&!@$VH|LM@a};)t=SzvREpb-Dz72Fdxz5-Rt{2ziuMd8PL+_W*t>%H7{wK;li;6Uj3c6@ps8IRoa@G8<9;gj$aDD-DPc@E!= zLVn_$2t610-Nb#t`_Yb=c5+;tiocogF#I`xZX^8){1N^J&&2OU_z1M$E8w>o>SJT%VI*xCNKkoG@$~hk9LGC^GPW53+I2K+4 z9Y8zstmDk_(g}Z_tKDy$Kc^#~3@!gM%5fa{cTBd6>v&h} zJb!H@UgK$m{1oVY7ou{N;OUY3$s;U0-M8!2QB{c8`hg8`AD?wejaW z(Hzn@Lhq@M;r~78J?b>tzaLx#*TQ#TPRczB4u{XdD=04;{zpUKgJbtG$60>u(Kp|H z;tJ%h6E!Kv@2F#9>br>aJJ9oZ1N;w$W1#Omg{Wssn2mDn$10@ng7#Bg(#E-g^w;nY zSehU0SKINm_Ca@W74m}5XulX5&rSWqz^c}Pn zcC+B+@MU-=_E$o`|DA-Ne$aE>Rw|es`=YQvbl*4}f0f~7a5nT_sekAFLdtuXa@-eh z!~Yui6#N4A!0s;SJO3Q~_&Xchu?Kdp>&Bm(_&mofCjBt{4*x^(>$rEnz7fAK!I{`? zhnul;J!_4g^YVVu-$Uc^oZ>#|IMTwW=v{}-eR~6o<}){Vb=y0$8ZaD zAF*C*vEK|`N5WZSj)>ednX&v<~@?qgR%+!LSSdmq5Qa>fd)2$Bpay z$e1`B4-@hGG<07oL|Q*rlXky$K3_}zPFR$F=hri%)#4!P4mgIPF9Jy{Dho6RU z0_=~y_l_kv+6Xs6&x^j}eM=mUL+kT6`cFXb0iFx|PUSl5yTVcU^L*|-svCZK!rS1N z&~tqU()P!tq_;xv1>PUei>db%^ge_?!EWd^COre%F0KPrY0u-R$BnQy`i{TzNWTJ) zjFERCy%HXQ-+^#2ya+x3>rvj-@EZ6&Y)U=OhW4B1XvfX@lv@wKXJhZY@I18uy=~Cn z%h>+?$-f9@$Dik^68L=-+D_BazXZBoI)3lN&;4*3_RmtUq1b&4T@RfH=b)bzx$E@V zq?bY0^#$l(hkx&Rt_Ks)p9f14R~>i)bpF)FUj_6}g^ug{NPBKFKGzM;bG}bqj=!eZ zIZvjL{uUODiO2KKBjh_C+$W6B^UB3Bdd2a3GIT#1hF&@1@;kEM3H`es*X^3v9}gdb zz4224KlYb)tI*p7v!WL}PTfDJpf?}Co}Wk1uQQ5eeB$i9LLV% zi_v%eYfSl`7gv+^y~6Q%0)B^Je-(TZI=)}Qz6kLwg1M`-<#s6IdYP^oyis&zaH=yr5 z$M2Kc*&d#2T*tn^ZWVDqp#gr4%Xxn(`KMv$zB(AY?8y7UZ(ttu{hg8HwlH$rbpvVl zFZJx7&d48u*WlNCw)>ppJa)aWirpsYJZeb&+ffh4_f+(5gqB+r{W9qL_c%q+^Ze74 z^dLBbc-=SK<7YF}eiHg~q5Jd)*k?iC?|es-o&=p2Kauu#KhAr9*H{bx^`Z0Q71G-% zuN?LC`?lk-6nd^tz60f@+_~7fzP3TX4_pe}e|u2gRj?-Zu4}fV>%?m8Z$jVmxWAv= zj=k%a=Tz7Io!EZ`U9Y^~PfDTIq-KZ^c`j; zY3H-^>redTLf#3kf-k_MiN|qgJf~ytIqog=TcK||SzdnhN5NiH;K3OGwwvp>=ZY)v z{~&Zf@^?Ghx!yW%t*74^yJCL^^!@A^(r-ZTL6eDdI&nMh>_^Wd`uP?6EY!<;ddZk} zsgA!Xa0h%7`_?hz!TLE5$D_X)w!)5Ll1?PQ1^gT)qvyGPHfhfX%_;YK9QpUE^ReFu zpMhV)o!EJvcnkX_$lri}LEF!9;`;0S^}Xc+;=_GX3GACf*8|sm+viE-M`Pa&Ho?*N z(C;&z$86`4*xT=SlHQ6R-&@O&KN;Qwzkrj_Z%lhQk47WE9&Uks&|gb@E3k7OH9>C( zoCnLI@4NO8?7AYK23Nq7DCa%w>~Gfz*PD&#d!Kh*b^mo8^qjr~eb?iml;ht^JWX7# z2d;CyvAY>A#Ljs>oj6-!-xFHTr$~d;^JNYEmBF9;%M{Wr@cR+8o%>_w`nrVl57>FW zDMfxQm<*pJ5XVyk^4r6X(DM_?q|W>w!{7qglX8bce+N{La=dR3C+&IK`^7B$`2F-k z{M4e}w!=~ARe(F;=g|AaMcDZs?6|gFu0r4Y^vjgz`_oo*jta@#1&M3Fz&Ho(l(2&U^49xB@?( zb04PsP0(|n^T~7HdH8qV@;gL3{H0*u20is1zbnZ<5qUk17B zlkra_e>{GDKk&oCVc0ugZ9msB$K{`tn-6`GssDbSak{Vjd$z3T4}y7#$98H?dJw!C z$DTi4B<*{x?@n8&&*#v2*8)F7q3?fH&^wnn?C*-`oee!Fub^D_u}a7thx%QHp6^b_ zQjY#dVRs!|21j9+9sLi9%X&0MZzB8=Zotm(6Mj!0gP;5HLo#UycHUoo2VaUpGWw1; z&pEbVLCRCl@BWUH6Y%f&bUk_rJKtx2q`X3uHyZlg`*#=0^ZTvs<-U9i_CqLdJDiWA_wX_Ju{{=Jw;eXZ4z;8KHX|y?EiwDh;txZ2`|MjNtc1tzhAi%|NdU1H;T?**E83#UFf@? z`0i|boQQlW{0Z9MB>(I8F|H5Zo9tiT0X)Y|$ARB(dQd-_A*nWU&p*B^IM3YQT{lnH z5IfIr&UgEB4SIEHpR$yj8@oyHPw08GQcQez4z#{nf zKH~o01pS53dHgi`j-OAEd;jf--w%i*J95`6->V)V|2p*Szh>m00cS(&^9<=v;aQaT zCcK;Y^ym4YC;A^k*Pr7l-#8q{^YP<&zbwYy`Sm{K*2K?M5cAaaq8|3Hw-wOKN`8B2 z`#7FFZ*@W6-+fp=&&@Tl{|Y`qT=r*i@`pk9tIuse{8u1w%# zC-i(e41fMk*M4>!6+*w10u0XmPcHZ0l-IncF1ig=;@1U>Yzcl`bA^!#T#_#2DDSGys?^*NFUjpwz z-~H5iXurEIZ9@Mp^!x=VL|2fd#2-{=#60|+)6VK_;`~UaEksJMf z$hX4D$o(Cu-y{7VX#1UtzWw|>Y2~+(cHMPyqz$d9ZWeZAH0t&|eDuj(G<9o^Kpq#_RZSz4RQP9lO|foHCTx9r|wH zlzLwoQ=aqk3G{C!zKt>U^ZS zjGXM$_*~~#ptqQEY+uW}3q#+nenWo)?Yb1X?>~1VzaRQ{D7N=N^xuZ=2cDl!z;8C> z_S^rqu$3+uIH}jNRK0a`?Cz?RfgxoUFiEA&ikSDu>T)K|4Voc zcK1NCQ{VBO_wF~goBPW3#J|~k;NVX92ON&vb1KzJeb4BG{3ci&1<$Llr;ej$=;bHw z6nF)^7kl^d;^@2Xcy6wUp5xkns!iP98%pDMIeZ*$hL@nf8+tBw-*rBDulIcFzYpO4 zJ_md60}F`zLduy3{XLld>wI?nxgM0opXdD(@pmaKfIt1@M&I?eK6+DWC+EW$^xQxH zBJKIb?+#~@zXNuoyn^^^M85Od@mYm(N1}f-tczU}=scW{pHJW@?5ksE|Jtr!qt_X| zeCXK@7a;#0`i?#oz58Jc;%J7R=O^#gui<|c`u^U=c{LTgEpRsaucN;WHl-ZrgYD+} z<+^J8UGaA@d;sRc&-<_x`IRZR2Xx%}Jv_(eVS_W;fx$L*ikx5mzW)$&&3 z!1nHseuJ3uE}*=dpzZAVdIUc|!S&cLL;q~@Jtx?X+sW^bp6As%a^WRB`^=}!9w<3QM`W{*sf39mS(Vqmbg0|y*q-{sf z!LFlz|8hOth`#qb&!6twM^V1#=!Z$Cp1)})fA{OSbUa>7eY{Ut@5PjBygNuw$G^X4 zF+9(KE6J#eU+)7i5Qq1s$IyEhJ%1PAxz=`dUgSi7EBe;|cH$|4|H)|5kxx;?2eibN}YNp=T zT_1~6z@;eIUY>{Czbm5WySDx8xGYQ@{U!d2QvMCFIP%)i^)i{X-#cDM-}UT773a?Y@~6OT*q;DhFYR~lSCx^U0`pVftoS*d`g#6#zjOT7M6V~zi@x_& z|BljjcD?W(bR2f)Q$EQg^?V2a0lz1s*8*NbydbFx4m|f}$Iid^zS{bcUkN@7T~F*s z<30z!ozW{q`Z4Ia>UZSJh|hcaYRdI{%rV$GuXm9C3VZJjWT$=?_uO+AdYieAcnU3f=Z#=j5z;A!}BdkU} zYC_KyRq*4wv4gbxi|fZi{9KIv5bAd;`tA$QpjQjMdth1gDhr+Ya^F1TBeAlHB z=&hh!kaPz5S3tk}4WwMx=M>7D4cDXZ`rH|Zo?oV*_cHtx#=b*0Bj0;vGs<#itfb{xtXz zY=rzG_#&jZBfzt<@5q6lrsR);xykpPBL%GpGPI#%mbM)Ts z`;Fsc0E!!->y!J{Px$vc=L-B=|GLOsKdvVI3$(u+AG^u_5U!WFE=b{q$bANnaPe0C${c)l)+UTxTvc*eq;pyST? zoOhG4djh@5_;+4yqdjUPKRw2u{opzqyUtyR|Ml=`=ze?}-KhG)eqVG6%9k6}p zAm0(A*O&AH=sU?d#9xE_i{Zo2bBpa?mHfxya@xc1KH16l-nWi)9m=^3CQ-l9Ft$G& zhkfyXId*rVe-wV3!SB&?+<4A*-Sr%NG`taah#0bneYqb?T|YkZ$|EWg#H&%fWKqtfS%*` zHq!oH%k|#+ZKgcmR|ZpV4(#T`HLwVB$J1Nne+a#=wk16pZiD)N7eCcw;&&Y6M&I%E zIceV|eJ99I`~%>4{5mdNcOF3QI#`KvQBVDyu=C+u1V>QrMED+b-SC}dEdJEDeT!1A z^Y{7~zvc1Q7y4bgC28kX9{fy%{_f54%A&so&PFdAad{qWg53Lmcl2(G$$y^wYQ$R{ zx$D_Gq`gOMjnQ*FtiVqhZ^FgUc5^)C z*FSOi{!2YD&vouJ?S{k_l!lhI>v_`s z_BG0(_@rf|edqfQfBt<%KIAoG;;Tn`8G&6$`a1YEei~5Fuh4z)KJ+{vj6xbvBKFBYJo?~)RF4aovf_xC1tbq1-4{k;7 zcc$^^d9QAY-0xQ9D97(Ij%Vx99{ug`7Ptfc0B@z7oTOcktXFU3uKR1zb3O4r|61~& zhKFG{f;fMI_S^O7RYFgH-WNT0j7EMY_M_o0SQ9^x=8;&E&6wKM;rOh~sJ@a_^aze~S6g?}=qed#;{Gx*Gl+=Z*vS z0q+$j$M|)exxVoa9M{rbyP@TJ-d;@pL$ERS_hSDbv|pa2 zVXWUU^ryqc=sW*!C;xr&cf(Qm@$Vz`*9847(Ehgnolo=db0@UFmXmhhbN<`k?(41( zj)xqSR|r;wjt|$F?BqKijnn#hPV0)E>-R&XPa#nKyC1zq+`hLIM86$$-DriL=WzS+ zW%PDh4*fZpbaw2Xl*k((ZwpUG?!90*`Om`Vkoyi{ziIEho{6643GZ=^C(qFh@#{F+ zf}QXF3+!0*%E3F}Y3MuOJ%9W6pS`K)XV8B3yniA2o}cvZyFqv4{oy^tlbiCJ!X)G$ zL;r5k?+dP<_V+!+?Y#1xdR~nGBKUV+4k3LLdai?g$+x{-M+%|$GIX8jgTDKq@6>*m zbDbD~p7-^yh`SQ`)uHX>z4BPfvtQbvXPkeL?u1@F(k z*L}@6eHU1XAN49w&T#xZfFJKap2zK9-wADZ$Ne1qTmx~F`n!JnWjXTc#PKBQg_Qpi zT#p0&v`6o3=soRp{HW)A_dHkLZW7jf9XdC0e2S0Y~xH$%rqGVyyZa{S*+{t+?d+s^(z^)K{5lKb#3;;4zh{oZ=d zMeg@Y*8%tQZ;(5F3Xs;`b>JQ9<37IwJ?FpQN&NR@no%FWA6`q^^|S=`?t7a_Po*BV zv*WiP^6A9k`dAEozc1cIn&MLb9hWuOAA=ku{Yd`H`0a?mb=Wu@caBT%_1<5-SKF^k z@ay{Al=N|w<9zQ$zWa&iE`RrCf7)J-H~Z@~D%ui1r$hfv#&U zhW?&wH3}PGDdZQyip0?y+TU-H_PkUDecv%XFXzOs_4RiFJJ5d$KONBX9P7EPupNqg z0DJ}eve-XE{si)GhI6rZUA+yx2FTs#FT~GW=sMuJ)B8*t{FTMeH0bwA&sEO1%aMPB zoxdA!U;GHU_s5IT^F7OXbQ5~BkxwPgF68^};knuGoLP~3emRzUJ08!(kN2`k=sPa9 zllC3L`7o0FrEne8AH^gMB>y#d3k%LwYOpJbwc6VZ`Bj=KS;=;=J#V{ZFtS z_MTgCliDB7g9_Bgdb;n-#m}wCk0tFmtc3q7;9%r8!=3Oqwdy?<@)p0Wvy`1>*9C>q@+ zkY5i$lJnH{_zcRS*wo)0pM)RdT8{&=Q{OY2U^oyz$~ThlxySX;c3y>i9d@1vKA;@e zL(dnsm-E1OaD85iKkt)^2;>_49f{m|=lbdQDd+c}=zT+6tMTLf9fQ9&;e7NzhrZJ+ z#m;t^h`-C=Eyyp!Pd4ml!~W={z#`M{#WSlMRFp47GjqA?;N<_ zzK+0oI~YHHM;%7Fz6ZOWxITKH^1ko-S_nU$o2`F+@-Kw0e`it7SLA<4+I9Cm;`F@W z`@QF~KGbsxY=)ovU;*sCx2`4s9Ln*0^%Z)qt6w1>K>k|j_iyK`^*M$vMmrf-Mf_}sGm*b&N09G2<$Cr7dR>se0HCPUDfP7cqp@>; z%YxkVQWNAY@Y5bYz8jy6yf2&wPeU&oY5VyH@;(2oqPz#7_c_PoiRkx4ZoG~I*E#P; zp3AJy;n=N!*4z5Iua2jj8pOLA`kv)EWGMOHLiZitEiAt;@{+Wh`;h%+{Vqj5g7`q{ z?^v!uzK-~skv53Y;W%g z=OKR_cEtZx(0cez(2o3{q36?$*q0{$d*QdxINn3=B&=Q#ApPRd2PHs$-hsVNFiB6r{L9oF$; zIj^Aadi?@%x~@+`?(YwuNA7!t>$m&XD#~4ip6gd>{P-T@I^_2n_iMj%8i(s(LF{+K zA=KOQ%aiZ9qa$$~MSe}v{*Iv__PdEQH+r5QJqMOVZaL0}htT)k^j^}hQI6}E_hIc_ zS8qk%`_o4;aroY99G+v|MGqu7o?VX?;Me^u2Wi(;?|rU&k0O7ZxaY*gQH}fu;Umah ze_ZEoLGE{!qWHZFZlJwghaFc>AoqOk_lTnSvwYWsUF3U?yB7Phf$A*4Nb zc+OjezTb-$Q$N=m?~%@@SCPA}4WV4Wr}&<0KewV>_ix`%I=0P`^Qhb>-Iug?q1y95 z9orthXQ=12S+UxSynEg}`}V0ZX2|GT{rjZ!sa2;{vn)wx{yz%;>;Fl3|EtbFQTT7I zccOiHu(eYt-cWv`?U7mf2U@%bT0EK64^ibDYL}=MIymG9TR$Fb@g5wb;NXukv&Kbi zIkD;e-IAGIp7jZp8_L%%#I)Ll{2lE1X>E`6+NIV0U~lh3@gHh|jA()M&Y$$!hw2dP zH#Qw>7bTC?k4k5>Kg-)+elyFzzH$9q8E?Bp>BY)p)BlwAwO%sOdPUhq$rCLvz4`}B zJO^5w>23FayZyhlKmV=uPISBU{1^4Rp>Rlkfb*LulB{YTkF z$rCLvz4`}BJO^5w>23Fas(oxfh0^?aeca6u@)NIvSbOba^P|$*rIoS%)9UA+Zg;S? zQz+h0zWo;^PgF0h`suaPKBF0L`($*vf0w^_>l?4$-&K!)dbt_Z{u$*bv)VhY@ut-- zv+DJ4D<`A-BcuI=`XjcS*mSI2lsr~HDxJ~(qRLYqDmRp`U5IJ53;ByTKhbg%)sHuh zSiQ8~N75Rfb}}}ej811(IT;;yTKyep^*zwyNo%_vXm;ss*Yw(l+9P&7i7F>b9x5l4 zpV9s-C(3WEJT|SL5EIo8`N`<~sJPNA&**snDgL6`HKXfkc{08JGs;hTB*KfRf@s<iwVgt{ygy zA_{LnaUv{&l>#Y9ftG^U=ex5L3t1o+ZY)LD(9l?C`LGTgJDg7-kWxi7NEAqbsOS-* z3L>PVqe6lLghVJij;J_q{i7W;S*?=gq#9KlFYi>D3i?o$pz%&vza+ylp&IKR6Hj$b&d6;;>R2DCB6;EJjsux!}_etzVEzk z`aE8h&-36s(m3L<%!AM4>;spq*#Yho`zAWvgD%@p#aKF_`6V9elEy>hRq3Lhs`Zga zva#~dN(XshrFlv5P+rw~vFnoRtt(#j{3LxTURNK^N?zB#*l;{v_WPu&`onel<^dlz zlrQmob$#)^d2RX_hn4hgaT_1wu&l>%_8Bkvp*-kGd?^lIsP)zL#fR1<@$~1X-~0I2 z@3K1zp)U8&Pj62?KRy2#cg4GK@%>-E{6l~6y1pKzvuS(|duMxVBYQRUh5MTN(D%tp z^7G*g^-7j?fP*f4;3OY-#&KNoSrwft=D_4%%cO_%ot z*A>g>JnSP6>Xqh+ao}Ao=L^+w%?BLkmEyK~FfQpqH{=<#~+IB>AQGmd@WLoNE@ zs*Ae7^UVt$ti*S%N8%%ndg0HDmHKnm=38|=+z0k;>+OrT=~X?BdU1Z>`{o7Tm&bU{ z&${e0ALFdUI9DHfu6S1*>vFxk&)K)-VSO7u);Q0Db)~w%v(1yo!Q(v1pLQC3m{s4f z9>@oaI(#kkqTLnyCmJk?wtm;X*l;|~e%<#s^r`H*$Gl*n>uPO0+c@VbJ0Ivt>tm}2`C*~wYo+|q zfxc7^$8CDhWgE(a4(fytTUHl3Y^8jRgYSz6FP7EM?SFpo#qF1$%QAgE>bd=S zNAo;}{ydWUX!g-{|0P`N+dThy5PCmp`ngE=!*s6a`CK>se5U%-`Ap}td4GSB&h31i z+yCqLhiMUz7~T5oTxwXUz; zXmxK~Ut3#ibvHWQ^^I1m+fi|4)YrPopFY^U`Mj!7U%j1klUw3N{bm8`_2B!5H@{U8 zeLXb|6@Jg>t&3j1{l+1a@BjJmu6`p0eLY>EW!aMY`7|dlzji#j)gSNX&0~$1-hT6U zU6H=77#bfIG)~6HqyE8QS9R^+=+@w9%rft7{&rn6 pYp3MJ{3}?cGkJZJf1!96bl~v*bU<&@-p+ZwMY2WpOT(tve*n)q-vIys literal 0 HcmV?d00001 From e747bb232ce65adb97938cba518cb366b108f36f Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Thu, 5 Feb 2026 13:14:18 -0600 Subject: [PATCH 10/19] update hermes --- ImputeFormer_ms.py | 538 +++++++++++++++++++++++++++++++++++++++++++++ hermes.py | 12 +- inc/test_ms.py | 136 ------------ 3 files changed, 541 insertions(+), 145 deletions(-) create mode 100644 ImputeFormer_ms.py delete mode 100644 inc/test_ms.py diff --git a/ImputeFormer_ms.py b/ImputeFormer_ms.py new file mode 100644 index 0000000..b211e33 --- /dev/null +++ b/ImputeFormer_ms.py @@ -0,0 +1,538 @@ +# ImputeFormer_ms.py +"""ImputeFormer_ms.py + +Baseline wrapper following the *gin_ms.py* CLI style and using the unified +*inc/test.py* Tester to report f1/nrmse. + +What this file provides +----------------------- +- A runnable multi-snapshot imputation baseline that takes a few observed + snapshots (times) and reconstructs the full diffusion history. +- Same tester contract as other baselines: model_fn returns y_pred with shape + [num_nodes, T] (times 0..T-1). The tester appends the final snapshot y[:, T] + automatically via test_fix_obs. + +Observation pattern +------------------- +- Use --obs_ts (or --obs_time alias) to specify observed time indices. + Example: --obs_ts "0,3,5" or --obs_time "5". +- Use -1 to represent T. +- If --obs_ts is not provided, use the last --obs_k snapshots ending at T. +- We ALWAYS include the final snapshot at time T as observed. + +About "official" ImputeFormer +----------------------------- +You asked to *lock* this wrapper to the official ImputeFormer implementation. +In this execution environment, outbound connections to GitHub raw assets are +blocked, so I cannot vendor the upstream source code here. + +Instead, this file ships a self-contained, ImputeFormer-*style* Transformer +imputer with **low-rank attention** (Linformer-like) to mimic the paper's +low-rank inductive bias. The interfaces (args, tensor shapes, tester contract) +are the important part for your ditto-ms integration. + +If you later provide (or vendor) the exact upstream class file, you can replace +`LockedImputeFormer` with the official class while keeping the training/eval +pipeline unchanged. +""" + +from __future__ import annotations + +import argparse +import math +from typing import List, Optional + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.optim as optim +from tqdm import trange + +from inc.diffus import diffus_gen, b_estim, SIR_STATES +from inc.test import Tester +from inc.utils import seed_all + + +# --------------------------------------------------------------------- +# Helpers: obs_time parsing (mirrors gin_ms.py; ALWAYS includes T) +# --------------------------------------------------------------------- + +def _parse_int_list(s: str) -> List[int]: + """Parse comma/space separated ints. Example: '0, 3,5' -> [0,3,5].""" + if s is None: + return [] + s = s.replace(" ", ",") + parts = [p.strip() for p in s.split(",") if p.strip() != ""] + return [int(p) for p in parts] + + +def _resolve_obs_ts(obs_ts: Optional[List[int]], obs_k: int, T: int) -> List[int]: + """Resolve observed snapshot indices in [0, T] (inclusive). + + - If obs_ts is provided: use it (with -1 mapped to T), clamp into [0, T], unique+sorted. + - Else: use last obs_k snapshots ending at T. + """ + if obs_ts is not None: + ts: List[int] = [] + for t in obs_ts: + if t == -1: + t = T + t = max(0, min(int(t), T)) + ts.append(t) + return sorted(set(ts)) + + k = max(1, min(int(obs_k), T + 1)) + start = max(0, T - k + 1) + return list(range(start, T + 1)) + + +def _make_obs_time(args, T: int) -> List[int]: + """Convert CLI args into final obs_time list, ALWAYS including T.""" + obs_ts = _resolve_obs_ts(args.obs_ts, args.obs_k, T) + obs_time = sorted(set(obs_ts + [T])) + return obs_time + + +# --------------------------------------------------------------------- +# ImputeFormer-style model (low-rank attention) +# --------------------------------------------------------------------- + +class LowRankSelfAttention(nn.Module): + """Multi-head self-attention with low-rank projection over sequence length. + + This is a Linformer-like approximation that reduces O(L^2) attention to + O(L * r) where r=proj_k. + + Input/Output: + x: [B, L, d_model] + out: [B, L, d_model] + """ + + def __init__( + self, + d_model: int, + n_heads: int, + proj_k: int, + max_len: int, + dropout: float, + ) -> None: + super().__init__() + assert d_model % n_heads == 0, "d_model must be divisible by n_heads" + assert proj_k > 0, "proj_k must be > 0" + + self.d_model = d_model + self.n_heads = n_heads + self.d_head = d_model // n_heads + self.proj_k = proj_k + self.max_len = max_len + + self.qkv = nn.Linear(d_model, 3 * d_model) + self.out_proj = nn.Linear(d_model, d_model) + + # Project K/V along the sequence length dimension (L -> proj_k) + self.E_k = nn.Parameter(torch.randn(max_len, proj_k) / math.sqrt(proj_k)) + self.E_v = nn.Parameter(torch.randn(max_len, proj_k) / math.sqrt(proj_k)) + + self.attn_drop = nn.Dropout(dropout) + self.proj_drop = nn.Dropout(dropout) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + B, L, _ = x.shape + if L > self.max_len: + raise ValueError(f"Sequence length L={L} exceeds max_len={self.max_len}. Increase --max_len") + + qkv = self.qkv(x) # [B, L, 3*d] + q, k, v = qkv.chunk(3, dim=-1) + + # [B, heads, L, d_head] + q = q.view(B, L, self.n_heads, self.d_head).transpose(1, 2) + k = k.view(B, L, self.n_heads, self.d_head).transpose(1, 2) + v = v.view(B, L, self.n_heads, self.d_head).transpose(1, 2) + + # Merge heads for projection and matmul + # [B*heads, L, d_head] + q = q.reshape(B * self.n_heads, L, self.d_head) + k = k.reshape(B * self.n_heads, L, self.d_head) + v = v.reshape(B * self.n_heads, L, self.d_head) + + # Project K and V along length dimension: [B*heads, proj_k, d_head] + Ek = self.E_k[:L, :] # [L, proj_k] + Ev = self.E_v[:L, :] + k_proj = torch.einsum("bld,lk->bkd", k, Ek) + v_proj = torch.einsum("bld,lk->bkd", v, Ev) + + # Attention: [B*heads, L, proj_k] + attn = torch.einsum("bld,bkd->blk", q, k_proj) / math.sqrt(self.d_head) + attn = attn.softmax(dim=-1) + attn = self.attn_drop(attn) + + # Output: [B*heads, L, d_head] + out = torch.einsum("blk,bkd->bld", attn, v_proj) + + # Restore heads: [B, L, d_model] + out = out.view(B, self.n_heads, L, self.d_head).transpose(1, 2).reshape(B, L, self.d_model) + out = self.out_proj(out) + out = self.proj_drop(out) + return out + + +class ImputeFormerBlock(nn.Module): + """Transformer block with low-rank self-attention.""" + + def __init__( + self, + d_model: int, + n_heads: int, + proj_k: int, + max_len: int, + dropout: float, + ffn_mult: int, + ) -> None: + super().__init__() + self.norm1 = nn.LayerNorm(d_model) + self.attn = LowRankSelfAttention( + d_model=d_model, + n_heads=n_heads, + proj_k=proj_k, + max_len=max_len, + dropout=dropout, + ) + self.norm2 = nn.LayerNorm(d_model) + self.ffn = nn.Sequential( + nn.Linear(d_model, ffn_mult * d_model), + nn.GELU(), + nn.Dropout(dropout), + nn.Linear(ffn_mult * d_model, d_model), + nn.Dropout(dropout), + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + # Pre-norm for stability + x = x + self.attn(self.norm1(x)) + x = x + self.ffn(self.norm2(x)) + return x + + +class LockedImputeFormer(nn.Module): + """A self-contained ImputeFormer-style imputer. + + It consumes the *entire* timeline (length L=T+1) with missing values masked. + + Input: + x: [B, L, in_dim] where in_dim = n_cls + 1 + - first n_cls channels: one-hot value at observed times, 0 otherwise + - last channel: observed mask (1 observed, 0 missing) + + Output: + logits: [B, L, n_cls] + """ + + def __init__( + self, + in_dim: int, + n_cls: int, + d_model: int, + n_heads: int, + n_layers: int, + proj_k: int, + dropout: float, + ffn_mult: int, + max_len: int, + ) -> None: + super().__init__() + self.in_dim = in_dim + self.n_cls = n_cls + self.d_model = d_model + self.max_len = max_len + + self.in_proj = nn.Linear(in_dim, d_model) + self.pos_emb = nn.Embedding(max_len, d_model) + + self.blocks = nn.ModuleList( + [ + ImputeFormerBlock( + d_model=d_model, + n_heads=n_heads, + proj_k=proj_k, + max_len=max_len, + dropout=dropout, + ffn_mult=ffn_mult, + ) + for _ in range(n_layers) + ] + ) + + self.out_norm = nn.LayerNorm(d_model) + self.out_proj = nn.Linear(d_model, n_cls) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + B, L, _ = x.shape + if L > self.max_len: + raise ValueError(f"Sequence length L={L} exceeds max_len={self.max_len}. Increase --max_len") + + h = self.in_proj(x) + pos = torch.arange(L, device=x.device) + h = h + self.pos_emb(pos)[None, :, :] + + for blk in self.blocks: + h = blk(h) + + h = self.out_norm(h) + logits = self.out_proj(h) + return logits + + +# --------------------------------------------------------------------- +# Input building +# --------------------------------------------------------------------- + +def _build_imputeformer_input(y: torch.Tensor, obs_time: List[int], n_cls: int) -> torch.Tensor: + """Build model input from integer labels with a global-time observation mask. + + Args: + y: [B, L] long (0..n_cls-1) + obs_time: list of observed time indices in [0, L-1] + n_cls: number of classes + + Returns: + x: [B, L, n_cls+1] float + - x[..., :n_cls] = one-hot(y) * obs_mask + - x[..., n_cls] = obs_mask + """ + B, L = y.shape + device = y.device + + obs_mask = torch.zeros((L,), dtype=torch.bool, device=device) + obs_mask[obs_time] = True + + # One-hot encode and zero-out unobserved positions + x_val = F.one_hot(y.clamp(min=0, max=n_cls - 1), num_classes=n_cls).float() # [B, L, n_cls] + x_val = x_val * obs_mask.view(1, L, 1).float() + + # Add mask as an extra channel + x_mask = obs_mask.view(1, L, 1).float().expand(B, L, 1) + x = torch.cat([x_val, x_mask], dim=-1) + return x + + +# --------------------------------------------------------------------- +# Main model_fn used by Tester +# --------------------------------------------------------------------- + +args = None # set in __main__ + + +def _call_b_estim(data, args, obs_time: List[int]): + """Call b_estim with best-effort compatibility across branches.""" + try: + return b_estim(data, args, obs_time=obs_time) + except TypeError: + # Fallback: older signature b_estim(data, args) + # Try to inject args.obs_time (string) for compatibility. + try: + setattr(args, "obs_time", ",".join(map(str, obs_time))) + except Exception: + pass + return b_estim(data, args) + + +def imputeformer_run(data) -> torch.Tensor: + """Train on simulated diffusion sequences then impute the test sequence. + + Returns: + y_pred: [num_nodes, T] long + """ + global args + + device = args.device + T = int(data.T.item()) + L = T + 1 + n_nodes = int(data.num_nodes) + + # Determine #classes (SI:2, SIR:3) + n_cls = int(data.y.max().item()) + 1 + + obs_time = _make_obs_time(args, T) + + # Estimate diffusion parameters (used for synthetic training labels) + bpar = _call_b_estim(data, args, obs_time=obs_time) + + # Initial infected count for simulation + I0 = int((data.y[:, 0] == SIR_STATES.I).sum().item()) + + # Model + model = LockedImputeFormer( + in_dim=n_cls + 1, + n_cls=n_cls, + d_model=args.units, + n_heads=args.heads, + n_layers=args.layers, + proj_k=args.proj_k, + dropout=args.dropout, + ffn_mult=args.ffn_mult, + max_len=max(args.max_len, L), + ).to(device) + + optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay) + + # ---------------- + # Train on simulated diffusion sequences + # ---------------- + for _ in trange(1, args.epochs + 1, desc="train", leave=False): + model.train() + + # labels: [batch, n_nodes, L] + labels = diffus_gen( + T=T, + n_nodes=n_nodes, + edge_index=data.edge_index, + I0=I0, + n_samples=args.batch_size, + pI=bpar.pI, + pR=bpar.pR, + ).transpose(0, 2) + + # Subsample nodes to bound memory + node_batch = min(int(args.node_batch), n_nodes) + if node_batch < n_nodes: + idx = torch.randint(0, n_nodes, (node_batch,), device=labels.device) + else: + idx = torch.arange(n_nodes, device=labels.device) + + # Each (diffusion sample, node) is one training instance: y: [B2, L] + y = labels[:, idx, :].reshape(-1, L).long() + + x = _build_imputeformer_input(y, obs_time=obs_time, n_cls=n_cls) + logits = model(x) # [B2, L, n_cls] + + loss = F.cross_entropy( + logits[:, :T, :].reshape(-1, n_cls), + y[:, :T].reshape(-1), + ) + + optimizer.zero_grad(set_to_none=True) + loss.backward() + if args.grad_clip and args.grad_clip > 0: + nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) + optimizer.step() + + # ---------------- + # Inference on the test instance (only obs_time snapshots are revealed) + # ---------------- + model.eval() + + y_true = data.y[:, :L].long().to(device) # [n_nodes, L] + y_pred_full = torch.empty((n_nodes, L), dtype=torch.long, device=device) + + with torch.no_grad(): + bs = max(1, int(args.eval_node_batch)) + for s in range(0, n_nodes, bs): + e = min(n_nodes, s + bs) + y_chunk = y_true[s:e] # [b, L] + x_chunk = _build_imputeformer_input(y_chunk, obs_time=obs_time, n_cls=n_cls) + logits = model(x_chunk) # [b, L, n_cls] + pred = logits.argmax(dim=-1) # [b, L] + + # Enforce consistency on observed snapshots (except final; tester fixes it anyway) + for t in obs_time: + if t < T: + pred[:, t] = y_chunk[:, t] + + y_pred_full[s:e] = pred + + # Return only 0..T-1 (tester appends y[:, T]) + return y_pred_full[:, :T] + + +# --------------------------------------------------------------------- +# CLI +# --------------------------------------------------------------------- + +def get_args(): + parser = argparse.ArgumentParser() + + # IO / runtime + parser.add_argument("--dataset", type=str, required=True) + parser.add_argument("--seed", type=int, required=True) + parser.add_argument("--data_dir", type=str, required=True) + parser.add_argument("--output", type=str, required=True) + parser.add_argument("--device", type=torch.device, required=True) + + # Diffusion parameter estimation (same knobs as other baselines) + parser.add_argument("--b_pI0", type=float, required=True) + parser.add_argument("--b_pR0", type=float, required=True) + parser.add_argument("--b_steps", type=int, required=True) + parser.add_argument("--b_lr", type=float, required=True) + parser.add_argument("--b_pImax", type=float, default=1.0) + + # Multi-snapshot observation settings + parser.add_argument( + "--obs_ts", + type=str, + default=None, + help='comma-separated observed time indices. Use -1 for T. Example: "0,3,5"', + ) + parser.add_argument( + "--obs_time", + type=str, + default=None, + help="alias of --obs_ts (kept for compatibility with other scripts).", + ) + parser.add_argument( + "--obs_k", + type=int, + default=1, + help="if obs_ts is None, use last k snapshots ending at T (default 1: final-only).", + ) + + # Model / training hyperparameters + parser.add_argument("--lr", type=float, default=1e-3) + parser.add_argument("--weight_decay", type=float, default=0.0) + parser.add_argument("--epochs", type=int, default=200) + parser.add_argument("--batch_size", type=int, default=16) + + parser.add_argument("--units", type=int, default=64, help="Transformer hidden size") + parser.add_argument("--heads", type=int, default=4, help="#attention heads") + parser.add_argument("--layers", type=int, default=4, help="#Transformer blocks") + parser.add_argument("--proj_k", type=int, default=16, help="low-rank projection size (Linformer k)") + parser.add_argument("--dropout", type=float, default=0.1) + parser.add_argument("--ffn_mult", type=int, default=4) + parser.add_argument("--max_len", type=int, default=64, help="max T+1") + + parser.add_argument("--grad_clip", type=float, default=1.0) + + # For large graphs + parser.add_argument( + "--node_batch", + type=int, + default=512, + help="#nodes sampled per training epoch (each node is a sequence instance)", + ) + parser.add_argument( + "--eval_node_batch", + type=int, + default=2048, + help="#nodes per forward pass during inference", + ) + + args = parser.parse_args() + + # obs_ts/obs_time normalization + obs_s = args.obs_ts if args.obs_ts is not None else args.obs_time + if obs_s is not None and len(str(obs_s).strip()) > 0: + obs = _parse_int_list(str(obs_s)) + args.obs_ts = sorted(set(obs)) + # Note: keep args.obs_time as-is; _make_obs_time uses args.obs_ts + else: + args.obs_ts = None + + return args + + +if __name__ == "__main__": + args = get_args() + seed_all(args.seed) + + tester = Tester(args.data_dir, args.device, imputeformer_run) + tester.test([args.dataset], rep=1) + tester.save(args.output) diff --git a/hermes.py b/hermes.py index 3f395bb..fc8e47c 100644 --- a/hermes.py +++ b/hermes.py @@ -346,12 +346,6 @@ def main(data): y_pred = y_pred[:, : data.T.item()].cummax(dim = 1).values return y_pred -def cli_main(): - global args - args = get_args() - tester = Tester(args.data_dir, args.device, main) - tester.test([args.dataset], seed=args.seed, rep=1) - tester.save(args.output) - -if __name__ == "__main__": - cli_main() +args = get_args() +tester = Tester(args.data_dir, args.device, main) +tester.test([args.dataset], seed = args.seed, rep = 1) diff --git a/inc/test_ms.py b/inc/test_ms.py deleted file mode 100644 index 61a1855..0000000 --- a/inc/test_ms.py +++ /dev/null @@ -1,136 +0,0 @@ -import functools as fnt -import numpy as np -import torch -from sklearn import metrics as skm - -from inc.data import * - -@torch.no_grad() -def _get_obs_mask(data): - T = int(data.T.item()) - n_nodes = int(data.num_nodes) - dev = data.y.device - for name in ('obs_mask', 'obs_masks'): - if hasattr(data, name): - m = getattr(data, name) - if isinstance(m, torch.Tensor): - mask = m - else: - mask = torch.as_tensor(m) - if mask.dim() == 1: # (T+1,) -> (nodes, T+1) - mask = mask.view(1, -1).expand(n_nodes, T + 1) - mask = mask.to(device=dev, dtype=torch.bool) - return mask - if hasattr(data, 'obs_ts') and getattr(data, 'obs_ts') is not None: - ts = getattr(data, 'obs_ts') - if not isinstance(ts, torch.Tensor): - ts = torch.as_tensor(ts, dtype=torch.long, device=dev) - mask = torch.zeros(n_nodes, T + 1, dtype=torch.bool, device=dev) - ts = ts.clamp_(0, T) - if ts.numel() > 0: - mask[:, ts.unique()] = True - return mask - mask = torch.zeros(n_nodes, T + 1, dtype=torch.bool, device=dev) - mask[:, T] = True - return mask - -@torch.no_grad() -def test_fix_obs(data, y_pred): - T = int(data.T.item()) - dev = data.y.device - if y_pred.size(1) == T: - y_pred = torch.cat([y_pred, data.y[:, -1:]], dim=1) - elif y_pred.size(1) != T + 1: - y_pred = y_pred[:, : T + 1] - - y_pred = y_pred.to(device=dev, dtype=data.y.dtype) - obs_mask = _get_obs_mask(data) # (nodes, T+1) - return torch.where(obs_mask, data.y, y_pred) - -@torch.no_grad() -def test_skm(skm_fn, data, y_pred, **kwargs): - y_fixed = test_fix_obs(data, y_pred) - obs_mask = _get_obs_mask(data) - unobs = ~obs_mask - y_true = torch2np(data.y[unobs]) - y_pred = torch2np(y_fixed[unobs]) - return float(skm_fn(y_true, y_pred, **kwargs)) - -@torch.no_grad() -def test_nrmse(data, y_pred): - y_pred = test_fix_obs(data, y_pred) - tI_pred = data_make_t(y_pred, SIR_STATES.I, dim = -1) - mse = skm.mean_squared_error(torch2np(data.tI), torch2np(tI_pred)) - if hasattr(data, 'tR'): - tR_pred = data_make_t(y_pred, SIR_STATES.R, dim = -1) - mseR = skm.mean_squared_error(torch2np(data.tR), torch2np(tR_pred)) - mse = (mse + mseR) / 2. - nrmse = np.sqrt(mse) / (data.T.item() + 1) - return float(nrmse) - - -TEST_METRICS = { - None: lambda data, y_pred: test_fix_obs(data, y_pred).tolist(), - 'acc': fnt.partial(test_skm, skm.accuracy_score), - 'prc': fnt.partial(test_skm, skm.precision_score, average='macro', zero_division=0), - 'rec': fnt.partial(test_skm, skm.recall_score, average='macro', zero_division=0), - 'f1': fnt.partial(test_skm, skm.f1_score, average='macro', zero_division=0), - 'nrmse': test_nrmse, -} - -class Tester: - def __init__(self, data_dir, device, model_fn): - self.data_dir = data_dir - self.device = device - self.model_fn = model_fn - self.res = dict() - - def test_once(self, dataset, seed=None): - data = data_load(dataset, self.data_dir, self.device) - if seed is not None: - seed_all(seed) - y_pred = self.model_fn(data) - if dataset not in self.res: - self.res[dataset] = dict() - res = self.res[dataset] - for metric, fn in TEST_METRICS.items(): - if metric not in res: - res[metric] = list() - self.res[dataset][metric].append(fn(data, y_pred)) - - def test_dataset(self, dataset, seed=None, rep=5, verbose=True): - for i in range(rep): - if verbose: - print(f'[{dataset} #{i}]', flush=True) - self.test_once(dataset, seed=None if seed is None else (seed ^ i)) - if verbose: - print( - f'[{dataset} #{i}]', - ', '.join([ - f'{metric}={scores[-1]:.4f}' - for metric, scores in self.res[dataset].items() - if metric is not None - ]), - flush=True - ) - - def test(self, datasets=None, seed=None, **kwargs): - if datasets is None: - datasets = DATASETS.keys() - for dataset in datasets: - self.test_dataset(dataset=dataset, seed=seed, **kwargs) - - def print(self, brief=True): - for dataset, metrics in self.res.items(): - for metric, scores in metrics.items(): - if metric is not None: - print(f'{dataset} {metric}:', end='') - if brief: - print(f' {np.mean(scores):.4f} ({np.std(scores):.4f})') - else: - for score in scores: - print(f' {score:.4f}', end='') - print('') - - def save(self, f, verbose=True): - return torch.save(self.res, f) From 241161fe541931a51dfbe6af13c64b2195d7ce4c Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Thu, 5 Feb 2026 22:13:56 -0600 Subject: [PATCH 11/19] update exp3 --- ditto_ms.py | 843 ------------------------------- experiments/exp3_beta_scatter.py | 312 ++++++++++++ 2 files changed, 312 insertions(+), 843 deletions(-) delete mode 100644 ditto_ms.py create mode 100644 experiments/exp3_beta_scatter.py diff --git a/ditto_ms.py b/ditto_ms.py deleted file mode 100644 index 9ec7c1b..0000000 --- a/ditto_ms.py +++ /dev/null @@ -1,843 +0,0 @@ -from inc.diffus import * -from inc.nn import * -from inc.test import * - -def get_args(): - parser = argparse.ArgumentParser() - parser.add_argument('--dataset', type = str, help = 'dataset name') - parser.add_argument('--seed', type = int, help = 'random seed') - parser.add_argument('--data_dir', type = str, help = 'dataset folder') - parser.add_argument('--output', type = str, help = 'output file name') - parser.add_argument('--device', type = torch.device, help = 'torch device') - parser.add_argument('--b_pI0', type = float, help = 'initial infection rate in diffusion parameter estimation') - parser.add_argument('--b_pR0', type = float, help = 'initial recovery rate in diffusion parameter estimation') - parser.add_argument('--b_steps', type = int, help = 'optimization steps in diffusion parameter estimation') - parser.add_argument('--b_lr', type = float, help = 'learning rate in diffusion parameter estimation') - parser.add_argument('--q_steps', type = int, help = 'training steps for the proposal model') - parser.add_argument('--q_lr', type = float, help = 'learning rate for the proposal model') - parser.add_argument('--q_hid', type = int, help = 'hidden size of the proposal model') - parser.add_argument('--q_gnn', type = int, help = 'number of layers of the GNN in the proposal model') - parser.add_argument('--q_mlp', type = int, help = 'number of layers of the MLP in the proposal model') - parser.add_argument('--q_samples', type = int, help = 'sample size to estimate the loss function of the proposal model') - parser.add_argument('--q_zlim', type = int, help = 'a hyperparameter to stablize gradient') - parser.add_argument('--p_coef', type = float, help = 'the coefficient gamma in the initial distribution P[y_0]') - parser.add_argument('--t_samples', type = int, help = 'MCMC sample size') - parser.add_argument('--t_steps', type = int, help = 'MCMC steps') - parser.add_argument('--t_keep', type = float, help = 'moving average in MCMC') - parser.add_argument('--obs_time', type = str, default = '', help = 'extra observed snapshot times, comma-separated, e.g., 5,7,9') - # Multi-snapshot proposal diagnostics/safety. - # The backward segment sampler uses rejection sampling under hard constraints. - # If a segment is infeasible (empty support) or the proposal assigns vanishing - # probability to feasible states, the rejection loop can otherwise run forever. - parser.add_argument('--ms_max_rounds', type=int, default=2048, - help='max rejection rounds per backward step in multi-snapshot sampling') - args = parser.parse_args() - return args - -class QNet(nn.Module): - @classmethod - def make(cls, data, args): - return cls( - eidx = data.edge_index, - T = data.T.item(), - hid = args.q_hid, - gnn = args.q_gnn, - mlp = args.q_mlp, - n_nodes = data.num_nodes, - zlim = args.q_zlim, - ms_max_rounds = getattr(args, 'ms_max_rounds', 2048), - ).to(args.device) - def __init__(self, eidx, T, hid, gnn, mlp, n_nodes, zlim, ms_max_rounds=2048): - super().__init__() - self.eidx = eidx - self.device = self.eidx.device - self.n_nodes = n_nodes - self.n_inf = self.n_nodes + 2 - self.n_edges = self.eidx.size(dim = 1) - self.zlim = zlim - # Maximum rejection rounds per backward step in multi-snapshot segment sampling. - # This avoids infinite loops when a segment has empty/tiny support under hard constraints. - self.ms_max_rounds = int(ms_max_rounds) - self.T = T - self.hid = int(hid) - self.gnn_dep = int(gnn) - self.mlp_dep = int(mlp) - self.w = nn.Parameter(data = torch.randn((self.n_edges, self.hid), dtype = torch.float32, device = self.device), requires_grad = True) - self.gnn = GNN(v_in = 1, e_in = self.hid, hid = self.hid, dep = self.gnn_dep) - self.mlp = MLP([self.hid] * self.mlp_dep + [2 * self.T]) - self.rem = (pyg.utils.degree(self.eidx[1], num_nodes = self.n_nodes).long().unsqueeze(dim = 1) + 1).detach().clone() # (nodes, 1) - self.neighbs = [[] for u in range(self.n_nodes)] - for i in range(self.n_edges): - self.neighbs[self.eidx[0, i].item()].append(self.eidx[1, i].item()) - for u in range(self.n_nodes): - self.neighbs[u] = torch.tensor(self.neighbs[u], dtype = torch.long, device = self.device) - self.adj = torch.sparse_coo_tensor( - indices = torch.stack([self.eidx[1], self.eidx[0]], dim = 0), - values = torch.ones(self.n_edges, dtype = torch.float32, device = self.device), - size = (self.n_nodes, self.n_nodes), - ).coalesce() - self.zero = torch.tensor(0., dtype = torch.float, device = self.device) - def clamp_z(self, z): - return z.clamp(-self.zlim, self.zlim) - def forward(self, y, orig = False): # y: (nodes, samples) - n_nodes, n_samples = y.size() - y = y.T.reshape((-1, 1)) # (samples*nodes, 1) - eidx = (self.eidx.unsqueeze(dim = 1) + n_nodes * torch.arange(n_samples, dtype = torch.long, device = y.device).unsqueeze(dim = -1)).reshape((2, -1)) # (2, samples*edges) - w = self.w.repeat(n_samples, 1) # (samples, hid) - z, e = self.gnn(y.float(), eidx, w) - z = self.mlp(z) # (samples*nodes, 2*T) - z = z.T.reshape((2 * self.T, n_samples, -1)) # (2*T, samples, nodes) - zI, zR = z[: self.T], z[self.T :] # (T, samples, nodes) - zI, zR = zI.transpose(1, 2), zR.transpose(1, 2) # (T, nodes, samples) - if orig: - return zI, zR, self.clamp_z(zI), self.clamp_z(zR) - else: - return self.clamp_z(zI), self.clamp_z(zR) - def lik(self, Y): # Y: (T+1, nodes, samples) - n_samples = Y.size(dim = 2) - zI0, zR0, zI, zR = self.forward(Y[-1], orig = True) # (T, nodes, samples) - zI = zI.clone().detach().requires_grad_(True); zI.retain_grad() - zR = zR.clone().detach().requires_grad_(True); zR.retain_grad() - # R->I - qR = torch.sigmoid(zR) # (T, nodes, samples) # prob of R->I - lR1 = torch_log(qR) # (T, nodes, samples) - lR0 = torch_log(1. - qR) # (T, nodes, samples) - with torch.no_grad(): - mskR = (Y[1 :] == SIR_STATES.R) # (T, nodes, samples) - trsR = (Y[: -1] != SIR_STATES.R) # (T, nodes, samples) - # I->S - zI_, uid = zI.sort(dim = 1, descending = True) # (T, nodes, samples) - qI = torch.sigmoid(zI_) # (T, nodes, samples) # prob of I->S - lI1 = torch_log(qI) # (T, nodes, samples) - lI0 = torch_log(1. - qI) # (T, nodes, samples) - with torch.no_grad(): - mskI = ((Y[1 :] >= SIR_STATES.I) & (Y[: -1] <= SIR_STATES.I)).flatten() # (T * nodes * samples) - trsI = ((Y[: -1] != SIR_STATES.I)).flatten() # (T * nodes * samples) - rem = torch.where(mskI, self.rem.expand(self.T, -1, n_samples).flatten(), self.n_inf) # (T * nodes * samples) - ptr = torch.arange(self.T, dtype = torch.long, device = self.device).unsqueeze(dim = 1) * self.n_nodes # (T, 1) - for i in range(uid.size(dim = 1)): - uidi = (ptr + uid[:, i]).flatten() * n_samples # (T * samples) - mski = mskI[uidi] # (T * samples) - if mski.max(): - trsi = trsI[uidi] # (T * samples) - vids, degi = [], [0] - for t in range(self.T): - for j in range(n_samples): - u = uid[t, i, j] - vids.append((t * self.n_nodes + u.unsqueeze(dim = 0)) * n_samples) - vid = self.neighbs[u.item()] - vids.append((t * self.n_nodes + vid) * n_samples) - degi.append(vid.size(dim = 0) + 1) - vids = torch.cat(vids, dim = 0) # (T * sum neighbs) - degi = torch.tensor(degi, dtype = torch.long, device = self.device) # (1 + T * samples) - indptr = degi.cumsum(dim = 0) # (1 + T * samples) - degi = degi[1 :] # (T * samples) - rems = rem.flatten()[vids] # (T * sum neighbs) - opti = (pysc.segment_min_csr(src = rems, indptr = indptr)[0] > 1) # (T * samples) - rem.flatten()[vids] = torch.where(mski.repeat_interleave(repeats = degi), torch.where(trsi.repeat_interleave(repeats = degi), rems - 1, self.n_inf), rems) # (T * sum neighbs) - mskI[uidi] &= opti # (T * samples) - # likR + likI - lik = ( - torch.where(mskR, torch.where(trsR, lR1, lR0), self.zero).reshape(-1, n_samples) - + torch.where( - mskI.reshape(-1, n_samples), - torch.where( - trsI.reshape(-1, n_samples), - lI1.reshape(-1, n_samples), - lI0.reshape(-1, n_samples), - ), - self.zero, - ) - ).sum(dim = 0) # (samples,) - return lik, zI0, zR0, zI, zR # (samples,) - - def _lik_seg_local(self, Y, zI, zR, TL, TR): - """ - Segment log-prob under the *original DITTO local-support* backward sampler. - This is essentially the original `lik()` but restricted to t in [TL, TR-1]. - - Parameters - ---------- - Y : LongTensor, (T+1, nodes, samples) - zI, zR : FloatTensor, (T, nodes, samples) (already clamped / combined) - TL, TR : int segment endpoints (TL < TR) - """ - n_samples = Y.size(dim=2) - L = int(TR - TL) - if L <= 0: - return torch.zeros(n_samples, dtype=torch.float32, device=self.device) - - # Slice the segment (TL..TR) - Yseg = Y[TL: TR + 1] # (L+1, nodes, samples) - zIseg = zI[TL: TR] # (L, nodes, samples) - zRseg = zR[TL: TR] # (L, nodes, samples) - - # ------------------------- - # R -> I (backward) - # ------------------------- - qR = torch.sigmoid(zRseg) # (L, nodes, samples) - lR1 = torch_log(qR) - lR0 = torch_log(1.0 - qR) - with torch.no_grad(): - mskR = (Yseg[1:] == SIR_STATES.R) # (L, nodes, samples) - trsR = (Yseg[:-1] != SIR_STATES.R) # (L, nodes, samples) - - # ------------------------- - # I -> S (backward) with DITTO ordering trick - # ------------------------- - zI_, uid = zIseg.sort(dim=1, descending=True) # (L, nodes, samples) - qI = torch.sigmoid(zI_) # (L, nodes, samples) - lI1 = torch_log(qI) - lI0 = torch_log(1.0 - qI) - - with torch.no_grad(): - # NOTE: this follows your existing `lik()` implementation style. - mskI = ((Yseg[1:] >= SIR_STATES.I) & (Yseg[:-1] <= SIR_STATES.I)).flatten() - trsI = ((Yseg[:-1] != SIR_STATES.I)).flatten() - - rem = torch.where(mskI, self.rem.expand(L, -1, n_samples).flatten(), self.n_inf) - ptr = torch.arange(L, dtype=torch.long, device=self.device).unsqueeze(dim=1) * self.n_nodes # (L,1) - - for i in range(uid.size(dim=1)): - uidi = (ptr + uid[:, i]).flatten() * n_samples # (L*samples) - mski = mskI[uidi] - if mski.max(): - trsi = trsI[uidi] - - vids, degi = [], [0] - for t in range(L): - for j in range(n_samples): - u = uid[t, i, j] - vids.append((t * self.n_nodes + u.unsqueeze(dim=0)) * n_samples) - vid = self.neighbs[u.item()] - vids.append((t * self.n_nodes + vid) * n_samples) - degi.append(vid.size(dim=0) + 1) - - vids = torch.cat(vids, dim=0) - degi = torch.tensor(degi, dtype=torch.long, device=self.device) - indptr = degi.cumsum(dim=0) - degi = degi[1:] - - rems = rem.flatten()[vids] - opti = (pysc.segment_min_csr(src=rems, indptr=indptr)[0] > 1) - - rem.flatten()[vids] = torch.where( - mski.repeat_interleave(repeats=degi), - torch.where(trsi.repeat_interleave(repeats=degi), rems - 1, self.n_inf), - rems, - ) - mskI[uidi] &= opti - - lik = ( - torch.where(mskR, torch.where(trsR, lR1, lR0), self.zero).reshape(-1, n_samples) - + torch.where( - mskI.reshape(-1, n_samples), - torch.where( - trsI.reshape(-1, n_samples), - lI1.reshape(-1, n_samples), - lI0.reshape(-1, n_samples), - ), - self.zero, - ) - ).sum(dim=0) - - return lik - - def _lik_seg_a1(self, Y, zI, zR, TL, TR, eps=1e-6): - S, I, R = SIR_STATES.S, SIR_STATES.I, SIR_STATES.R - device = self.device - n_samples = Y.shape[2] - p_seg = self.zero.expand(n_samples).clone() - - - yL = Y[TL] - blocked = (yL == R) - src = (yL == I) - - # IMPORTANT: only t in (TL, TR) i.e. TL+1 ... TR-1 (consistent with _samp_seg) - for t in range(TL + 1, TR): - d = t - TL - y_t = Y[t] - y_next = Y[t + 1] - - qR = torch.sigmoid(zR[t]).clamp(eps, 1.0 - eps) # (n_nodes, 1) - p_add = (1.0 - torch.sigmoid(zI[t])).clamp(eps, 1.0 - eps) # (n_nodes, 1) - - pool = (y_next != S) & (~blocked) - A_true = (y_t != S) & pool - - # replay growth to compute log prob of A_true - A = src & pool - frontier = A.clone() - decided = A.clone() - - logp_A = torch.zeros(y_t.shape[1], device=device) - - for _ in range(d): - cand = (torch.sparse.mm(self.adj, frontier.float()) > 0) & pool & (~decided) - if not cand.any(): - break - - add = cand & A_true - logp_A += (add.float() * torch_log(p_add) + (cand & ~add).float() * torch_log(1 - p_add)).sum(0) - - A |= add - frontier = add - decided |= cand - - # if A_true contains nodes never reached by growth => prob 0 - missing = A_true & (~A) - if missing.any(): - # set those samples to -inf - bad = missing.any(dim=0) - logp_A[bad] = -float("inf") - - # candR likelihood - candR = A_true & (y_next == R) - isI = (y_t == I) - toI = candR & isI - toR = candR & (~isI) # should be R - - logp_R = (toI.float() * torch_log(qR) + toR.float() * torch_log(1 - qR)).sum(0) - - # if y_next==I and in A_true, must be I - badI = (A_true & (y_next == I) & (y_t != I)).any(dim=0) - if badI.any(): - logp_R[badI] = -float("inf") - - p_seg += (logp_A + logp_R) - - return p_seg - - def lik_ms(self, Y, obs_time): - """ - Multi-snapshot proposal likelihood matching `samp_ms()`. - """ - assert Y.size(dim=0) == self.T + 1, "lik_ms expects Y with shape (T+1, nodes, samples)" - n_samples = Y.size(dim=2) - - # sanitize obs_time: unique, within (0..T], and must include T - obs_time = sorted({int(t) for t in obs_time if 0 < int(t) <= self.T}) - if self.T not in obs_time: - obs_time.append(self.T) - K = len(obs_time) - - # Condition on all observed snapshots: concat along the "samples" dimension - y_cond = torch.cat([Y[t] for t in obs_time], dim=1) - - # Forward once for all conditioning blocks - zI0, zR0, zI, zR = self.forward(y_cond, orig=True) # (T, nodes, samples*K) - - # reshape to (T, nodes, K, samples) - zI0 = zI0.contiguous().view(self.T, self.n_nodes, K, n_samples) - zR0 = zR0.contiguous().view(self.T, self.n_nodes, K, n_samples) - zI = zI.contiguous().view(self.T, self.n_nodes, K, n_samples) - zR = zR.contiguous().view(self.T, self.n_nodes, K, n_samples) - - # detach+clamp leaf trick (same training mechanism as original lik()) - zI = zI.clone().detach().requires_grad_(True) - zI.retain_grad() - zR = zR.clone().detach().requires_grad_(True) - zR.retain_grad() - - lik = torch.zeros(n_samples, dtype=torch.float32, device=self.device) - - segL = [0] + obs_time[:-1] - segR = obs_time - - for i in range(K): - TL, TR = segL[i], segR[i] - - # logits conditioned on right endpoint snapshot y_TR (index i) - zI_R = zI[:, :, i, :] # (T, nodes, samples) - zR_R = zR[:, :, i, :] - - # combine left+right logits for TL>0 as in samp_ms() - if (TL > 0) and (K > 1): - zI_L = zI[:, :, i - 1, :] - zR_L = zR[:, :, i - 1, :] - zI_seg = self.clamp_z(zI_R + zI_L) - zR_seg = self.clamp_z(zR_R + zR_L) - else: - zI_seg = zI_R - zR_seg = zR_R - - if TL == 0: - lik = lik + self._lik_seg_local(Y, zI_seg, zR_seg, TL=TL, TR=TR) - else: - # _lik_seg_a1() only returns `lik` and does not accept return_ok / invalid_to_neg_inf - lik = lik + self._lik_seg_a1(Y, zI_seg, zR_seg, TL=TL, TR=TR) - - return lik, zI0, zR0, zI, zR - - @torch.no_grad() - def clamp_grad(self, z0, grad): - return torch.where(z0 < self.zlim, torch.where(z0 > -self.zlim, grad, F.relu(grad)), -F.relu(-grad)) - def backward(self, loss, zI0, zR0, zI, zR): - loss.backward() - z0 = torch.stack([zI0, zR0], dim = 0) - z0.backward(torch.stack([self.clamp_grad(zI0, zI.grad), self.clamp_grad(zR0, zR.grad)], dim = 0)) - @torch.no_grad() - def samp(self, y, zI, zR, n_samples, compute_lik = False): # y: (nodes,); zI, zR: (T, nodes, 1) - zI, uid = zI.sort(dim = 1, descending = True) # (T, nodes, 1) - uid = uid.squeeze(dim = 2) # (T, nodes) - qI = torch.sigmoid(zI) # (T, nodes, 1) # prob of I->S - xI = SIR_STATES.I - qI.expand(-1, -1, n_samples).bernoulli().long() # (T, nodes, samples) # 1 for I->S - lI = torch_log(torch.where(xI != SIR_STATES.I, qI, 1. - qI)) # (T, nodes, samples) - qR = torch.sigmoid(zR) # (T, nodes, 1) # prob of R->I - xR = SIR_STATES.R - qR.expand(-1, -1, n_samples).bernoulli().long() # (T, nodes, samples) # 1 for R->I - lR = torch_log(torch.where(xR != SIR_STATES.R, qR, 1. - qR)) # (T, nodes, samples) - y = y.unsqueeze(dim = 1).expand(-1, n_samples) # (nodes, samples) - Y = torch.empty(self.T, self.n_nodes, n_samples, dtype = torch.long, device = self.device) # (T, nodes, samples) - if compute_lik: - lik = self.zero - for t in range(self.T - 1, -1, -1): - # R->I - msk = (y == SIR_STATES.R) # (nodes, samples) - y = torch.where(msk, xR[t], y) # (nodes, samples) - if compute_lik: - lik = lik + torch.where(msk, lR[t], self.zero).sum(dim = 0) # (samples,) - # I->S - msk = (y == SIR_STATES.I) # (nodes, samples) - rem = torch.where(msk, self.rem, self.n_inf) # (nodes, samples) - for i, u in enumerate(uid[t]): - if msk[u].max(): - vid = self.neighbs[u.item()] # (neighbs,) - opt = (rem[u] > 1) & (rem[vid].min(dim = 0).values > 1) # (samples,) - msk_opt = msk[u] & opt - y[u] = torch.where(msk_opt, xI[t, i], y[u]) # (samples,) - trs = (y[u] != SIR_STATES.I) # (samples,) - rem[u] = torch.where(msk[u], torch.where(trs, rem[u] - 1, self.n_inf), rem[u]) # (samples,) - rem[vid] = torch.where(msk[u].unsqueeze(dim = 0), torch.where(trs.unsqueeze(dim = 0), rem[vid] - 1, self.n_inf), rem[vid]) # (neighbs, samples) - msk[u] = msk_opt - Y[t] = y - if compute_lik: - lik = lik + torch.where(msk[uid[t]], lI[t], self.zero).sum(dim = 0) # (samples,) - Y = Y.detach().clone() - if compute_lik: - lik = lik.detach().clone() - return Y, lik - else: - return Y - - @torch.no_grad() - def ext_ok(self, x, yL, t, TL): - d = t - TL - ok = (x >= yL.unsqueeze(dim=1)).all(dim=0) # (samples,) - if not ok.max(): - return ok - src = (yL == SIR_STATES.I).unsqueeze(dim=1) # (nodes, 1) - A = (x != SIR_STATES.S) & (yL != SIR_STATES.R).unsqueeze(dim=1) # (nodes, samples) - reach = src.expand(-1, x.size(dim=1)) & A # (nodes, samples) - for _ in range(d): - nbr = (torch.sparse.mm(self.adj, reach.float()) > 0) & A # (nodes, samples) - reach = reach | nbr - Iset = (x == SIR_STATES.I) & (yL != SIR_STATES.R).unsqueeze(dim=1) # (nodes, samples) - Rset = (x == SIR_STATES.R) & (yL != SIR_STATES.R).unsqueeze(dim=1) # (nodes, samples) - ok = ok & ~(Iset & ~reach).any(dim=0) - # NOTE: allow one-step S->R (infection + recovery within the same discrete step), - # so R-nodes at time t only need distance <= d (not d-1). - ok = ok & ~(Rset & ~reach).any(dim=0) - return ok - - @torch.no_grad() - def _samp_step(self, y, zI, uid, zR, t, compute_lik=False, yL=None): - """One backward step: sample y_t given y_{t+1}=y. - - This follows DITTO's original local-support (right-end feasibility) design via the - `msk/rem` mechanism. - - Multi-snapshot add-on (Fix 2): if a left endpoint snapshot yL (= y_{TL}) is provided, - we *hard clamp* the **left-monotonicity** necessary constraint directly inside the - sampler so we never generate y_t < yL. - - Concretely (S < I < R): - - Nodes with yL==R must stay R for all t>=TL => disable backward R->I. - - Nodes with yL==I must stay in {I,R} => disable backward I->S on those nodes. - - IMPORTANT: We enforce this by masking *sampling choices*, not by post-hoc overwriting, - so the returned `lik` remains the correct proposal log-probability. - - Parameters - ---------- - y : LongTensor, (nodes, samples) - Current snapshot y_{t+1} for a batch of samples. - yL : LongTensor or None, (nodes,) - Left observed snapshot y_{TL} (for monotonic clamping). If None, no clamping. - """ - n_samples = y.size(dim=1) - - # Left-monotonic hard constraints (segment-wise), if provided. - if yL is not None: - force_R = (yL == SIR_STATES.R).unsqueeze(dim=1) # (nodes, 1) - force_I = (yL == SIR_STATES.I) # (nodes,) - else: - force_R = None - force_I = None - if compute_lik: - lik = self.zero - # R->I - qR = torch.sigmoid(zR[t]) # (nodes, 1) - xR = SIR_STATES.R - qR.expand(-1, n_samples).bernoulli().long() # (nodes, samples) - if compute_lik: - lR = torch_log(torch.where(xR != SIR_STATES.R, qR, 1. - qR)) # (nodes, samples) - msk = (y == SIR_STATES.R) # (nodes, samples) - # Fix 2 (part 1): nodes already recovered at TL must remain R => do NOT sample R->I. - if force_R is not None: - msk = msk & (~force_R) - y = torch.where(msk, xR, y) - if compute_lik: - lik = lik + torch.where(msk, lR, self.zero).sum(dim=0) # (samples,) - # I->S - qI = torch.sigmoid(zI[t]) # (nodes, 1) (already sorted by zI) - xI = SIR_STATES.I - qI.expand(-1, n_samples).bernoulli().long() # (nodes, samples) - if compute_lik: - lI = torch_log(torch.where(xI != SIR_STATES.I, qI, 1. - qI)) # (nodes, samples) - msk = (y == SIR_STATES.I) # (nodes, samples) - rem = torch.where(msk, self.rem, self.n_inf) # (nodes, samples) - for i, u in enumerate(uid[t]): - if msk[u].max(): - vid = self.neighbs[u.item()] # (neighbs,) - opt = (rem[u] > 1) & (rem[vid].min(dim=0).values > 1) # (samples,) - # Fix 2 (part 2): nodes infected at TL must never go below I => do NOT sample I->S. - # Enforce by forcing `opt=False` (deterministic keep-I) for those nodes. - if force_I is not None: - opt = opt & (~force_I[u]) - msk_opt = msk[u] & opt - y[u] = torch.where(msk_opt, xI[i], y[u]) # (samples,) - trs = (y[u] != SIR_STATES.I) # (samples,) - rem[u] = torch.where(msk[u], torch.where(trs, rem[u] - 1, self.n_inf), rem[u]) - rem[vid] = torch.where(msk[u].unsqueeze(dim=0), - torch.where(trs.unsqueeze(dim=0), rem[vid] - 1, self.n_inf), rem[vid]) - msk[u] = msk_opt - if compute_lik: - lik = lik + torch.where(msk[uid[t]], lI, self.zero).sum(dim=0) # (samples,) - return y, lik - else: - return y - - def _samp_seg(self, yL, yR, zI_sorted, uid, zR, TL, TR, compute_lik=False): - S, I, R = self.S, self.I, self.R - n_samples = yR.shape[1] - device = self.device - - Y = torch.empty(TR - TL + 1, self.n_nodes, n_samples, dtype=torch.long, device=device) - lik = 0.0 if compute_lik else None - - # Clamp left endpoint - yL = yL.unsqueeze(dim=1).expand(-1, n_samples) - yR = yR.unsqueeze(dim=1).expand(-1, n_samples) - Y[0] = yL - Y[TR - TL] = yR - - # TL=0 : keep original DITTO step - if TL == 0: - y = yR - for t in range(TR - 1, 0, -1): - y, p_x = self._samp_step(y, zI_sorted[t], uid[t], zR[t], compute_lik) - Y[t] = y - if compute_lik: - lik += p_x - return Y, lik - - # ------------------------- - # TL > 0 : Route-B sampler - # ------------------------- - - # unsort zI for p_add lookup (keep your original logic) - zI_unsorted = torch.empty_like(zI_sorted).squeeze(dim=2) # (T, n_nodes) - for tt in range(self.T): - zI_unsorted[tt].scatter_(dim=0, index=uid[tt], src=zI_sorted[tt].squeeze(dim=1)) - - blocked = (yL == R) - src = (yL == I) - - for t in range(TR - 1, TL, -1): - d = t - TL - y_next = Y[t - TL + 1] # (n_nodes, n_samples) - - # probabilities - qR = torch.sigmoid(zR[t]) # (n_nodes, 1) - p_add = 1.0 - torch.sigmoid(zI_unsorted[t]).unsqueeze(1) # (n_nodes, 1) - # (optional) clamp to avoid exactly 0/1 - qR = qR.clamp(self.eps, 1.0 - self.eps) - p_add = p_add.clamp(self.eps, 1.0 - self.eps) - - pool = (y_next != S) & (~blocked) - - # ------------------------------------------------------------ - # Step 1: A_t generation WITHOUT forcing reach_{<=d-1} - # (outward growth from src for d hops) - # ------------------------------------------------------------ - A = src & pool - frontier = A.clone() - decided = A.clone() # nodes whose add/not-add decision has been made - - if compute_lik: - logp_A = torch.zeros(n_samples, device=device) - - for _ in range(d): - cand = (torch.sparse.mm(self.adj, frontier.float()) > 0) & pool & (~decided) - if not cand.any(): - break - - u = torch.rand(self.n_nodes, n_samples, device=device) - add = cand & (u <= p_add) # Bernoulli(p_add) - if compute_lik: - logp_A += (add.float() * torch_log(p_add) + (cand & ~add).float() * torch_log(1 - p_add)).sum(0) - - A |= add - frontier = add - decided |= cand - - # ------------------------------------------------------------ - # Step 2: sample states in A, but DO NOT force all neighbors of new - # Enforce: each new has >=1 infected neighbor (only if otherwise 0-prob) - # ------------------------------------------------------------ - x_t = torch.full_like(y_next, S) - x_t[blocked] = R - - # forced I if y_next==I and in A - mI = A & (y_next == I) - x_t[mI] = I - - # candR nodes: y_next==R and in A => sample I/R via qR - candR = A & (y_next == R) - uR = torch.rand(self.n_nodes, n_samples, device=device) - toI = candR & (uR <= qR) - toR = candR & (~toI) - x_t[toI] = I - x_t[toR] = R - - if compute_lik: - logp_R = (toI.float() * torch_log(qR) + toR.float() * torch_log(1 - qR)).sum(0) - - # new nodes - new = pool & (~A) - - # ---- enforce infection-source constraint minimally ---- - # For each sample: if a new node has no infected neighbor, pick ONE neighbor in A and flip to I - # (only when otherwise impossible / zero-prob forward) - I_mask = (x_t == I) - neighI = (torch.sparse.mm(self.adj, I_mask.float()) > 0) - bad_new = new & (~neighI) # nodes that violate ">=1 infected neighbor" - if bad_new.any(): - # candidates that could be flipped to I: in A, and (y_next==I already I) OR (y_next==R and in candR) - # (y_next==I in A are already I; so only need consider A & (y_next==R) that are currently R) - fixable = A & (y_next == R) - - # For each bad_new node u, choose one neighbor v from fixable ∩ N(u) to flip to I - # If none exists, that sample is infeasible under model (posterior prob 0), so leave as is (will be rejected later if you have a checker) - neigh_fixable = (torch.sparse.mm(self.adj, fixable.float()) > 0) - # We flip per bad_new node by sampling one neighbor index. - # Implementation trick: do one pass "greedy-random" by picking first available neighbor per (u,sample) - # (this avoids heavy per-node loops and still gives each choice positive prob if you randomize tie-breaking) - # Here: random tie-breaking by multiplying adjacency mask with random noise. - adj_dense = self.adj.to_dense() # (n_nodes, n_nodes) maybe too big; if too big, replace with sparse gather kernels - # NOTE: if graph is large, do not materialize dense. In that case implement sparse neighbor sampling separately. - - # For minimal code change, keep dense only if feasible in your scale. - noise = torch.rand(self.n_nodes, self.n_nodes, device=device) - # candidates matrix: (u,v,sample) => u in bad_new, v in fixable neighbor - # build neighbor mask (u,v) then apply for each sample - nb_mask = (adj_dense > 0) - - for s in range(n_samples): - bad_u = bad_new[:, s].nonzero(as_tuple=False).flatten() - if bad_u.numel() == 0: - continue - fix_v = fixable[:, s] - # for each bad u, pick v maximizing noise among allowed neighbors - for u_node in bad_u.tolist(): - allowed = nb_mask[u_node] & fix_v - if allowed.any(): - v_idx = (noise[u_node] * allowed.float()).argmax().item() - x_t[v_idx, s] = I - - # optional: if you want strict soundness, you can assert-check again and resample/reject, - # but for minimal code change we just repair as above. - - Y[t - TL] = x_t - - if compute_lik: - lik += (logp_A + logp_R) - - return Y, lik - - @torch.no_grad() - def samp_ms(self, y, zI, zR, n_samples, obs_time, compute_lik=False): - - # y: (nodes, T+1) - obs_time = sorted(list(obs_time)) - segL = [0] + obs_time[:-1] - segR = obs_time - - # Allocate full history tensor (we store times 0..T-1; y_T is not stored here). - Y = torch.empty(self.T, self.n_nodes, n_samples, dtype=torch.long, device=self.device) - - # Hard constraints: directly fix observed snapshots (except the final y_T which is not in Y) - for t in obs_time: - if t < self.T: - Y[t] = y[:, t].unsqueeze(dim=1).expand(-1, n_samples) - - if compute_lik: - lik = self.zero # will become (samples,) after first addition - - # Helper: fetch logits conditioned on the i-th observed snapshot. - # If z has only one conditioning slice, reuse it for all segments. - def _pick(z, i): - # z: (T, nodes, K) or (T, nodes, 1) - if z.size(dim=2) == 1: - return z - return z[:, :, i: i + 1] - - # Sample segments in reverse order (right endpoint always known). - for i in range(len(segR) - 1, -1, -1): - TL, TR = segL[i], segR[i] - yR = y[:, TR] - - if TL > 0: - yL = y[:, TL] - start = TL + 1 # only fill (TL, TR) - else: - yL = None - start = TL - zI_R = _pick(zI, i) # conditioned on y_TR - zR_R = _pick(zR, i) - - if (yL is not None) and (zI.size(dim=2) > 1): - # TL corresponds to obs_time[i-1] - zI_L = _pick(zI, i - 1) - zR_L = _pick(zR, i - 1) - - # Combine evidence in logit space, then clamp. - zI_seg = self.clamp_z(zI_R + zI_L) - zR_seg = self.clamp_z(zR_R + zR_L) - else: - # first segment (TL=0) OR caller provided only one logit slice - zI_seg = zI_R - zR_seg = zR_R - - zI_sorted, uid = zI_seg.sort(dim=1, descending=True) # (T, nodes, 1) - uid = uid.squeeze(dim=2) # (T, nodes) - if compute_lik: - Y_seg, lik_seg = self._samp_seg( - yR, zI_sorted, uid, zR_seg, n_samples, TL, TR, yL=yL, compute_lik=True - ) - if yL is not None: - Y[start:TR] = Y_seg[1:] # skip the clamped y_TL - else: - Y[start:TR] = Y_seg - - lik = lik + lik_seg - else: - Y_seg = self._samp_seg(yR, zI_sorted, uid, zR_seg, n_samples, TL, TR, yL=yL, compute_lik=False) - if yL is not None: - Y[start:TR] = Y_seg[1:] - else: - Y[start:TR] = Y_seg - - if compute_lik: - return Y.detach().clone(), lik.detach().clone() - else: - return Y.detach().clone() - - -def q_loss(q_net, data, I0, bpar, n_samples, obs_time): - T = data.T.item() - n_nodes = data.num_nodes - Y = diffus_gen( - T=T, - n_nodes=n_nodes, - edge_index=data.edge_index, - I0=I0, - n_samples=n_samples, - pI=bpar.pI, - pR=bpar.pR, - ) # (T+1, nodes, samples) - - q_liks, zI0, zR0, zI, zR = q_net.lik_ms(Y=Y, obs_time=obs_time) - return -q_liks.mean(), zI0, zR0, zI, zR - -def q_train(data, bpar, args): - I0 = (data.y[:, 0] == 1).long().sum().item() - - # Keep training obs_time consistent with main()/t_mcmc. - obs_time = [int(t) for t in args.obs_time.split(',') if t] - obs_time.append(data.T.item()) - obs_time = sorted({t for t in obs_time if 0 < t <= data.T.item()}) - - q_net = QNet.make(data, args) - q_net.train() - opt = optim.AdamW(q_net.parameters(), lr=args.q_lr) - pbar = trange(1, args.q_steps + 1) - for step in pbar: - opt.zero_grad() - loss, zI0, zR0, zI, zR = q_loss(q_net, data, I0, bpar, args.q_samples, obs_time=obs_time) - pbar.set_description(f'[step={step}] loss={loss.item():.4f}') - q_net.backward(loss, zI0, zR0, zI, zR) - opt.step() - q_net.eval() - return q_net - -@torch.no_grad() -def t_mcmc(data, bpar, q_net, args, obs_time, keepdim=True): - - I0 = (data.y[:, 0] == 1).long().sum().item() - - obs_time = sorted(list(obs_time)) - y_obs = torch.stack([data.y[:, t] for t in obs_time], dim=1) # (nodes, K_obs) - zI, zR = q_net(y_obs) # (T, nodes, K_obs) - - X, lqX = q_net.samp_ms(data.y, zI, zR, args.t_samples, obs_time=obs_time, compute_lik=True) - lpX = diffus_liks(Y=X, edge_index=data.edge_index, I0=I0, coef=args.p_coef, pI=bpar.pI, pR=bpar.pR) - - tI_avg = data_make_t(X, SIR_STATES.I, dim=0).float().mean(dim=1, keepdim=keepdim) - tR_avg = data_make_t(X, SIR_STATES.R, dim=0).float().mean(dim=1, keepdim=keepdim) - - pbar = trange(1, args.t_steps + 1) - for step in pbar: - Y, lqY = q_net.samp_ms(data.y, zI, zR, args.t_samples, obs_time=obs_time, compute_lik=True) - lpY = diffus_liks(Y=Y, edge_index=data.edge_index, I0=I0, coef=args.p_coef, pI=bpar.pI, pR=bpar.pR) - - # Hastings acceptance - a = torch.rand(args.t_samples, device=args.device) <= torch.exp(lpY + lqX - lpX - lqY) - - X = torch.where(a, Y, X) - lqX = torch.where(a, lqY, lqX) - lpX = torch.where(a, lpY, lpX) - - tI = data_make_t(X, SIR_STATES.I, dim=0).float().mean(dim=1, keepdim=keepdim) - tR = data_make_t(X, SIR_STATES.R, dim=0).float().mean(dim=1, keepdim=keepdim) - tI_avg = args.t_keep * tI_avg + (1.0 - args.t_keep) * tI - tR_avg = args.t_keep * tR_avg + (1.0 - args.t_keep) * tR - - return tI_avg, tR_avg - -def main(data): - # parse obs times - obs_time = [int(t) for t in args.obs_time.split(',') if t] - obs_time.append(data.T.item()) - obs_time = sorted(obs_time) - # estimate diffusion parameters - bpar = b_estim(data, args) - print(f'[est] pI={bpar.pI:.4f}, pR={bpar.pR:.4f}', flush = True) - # train a proposal network - q_net = q_train(data, bpar, args) - # estimate transition times - tI, tR = t_mcmc(data, bpar, q_net, args, obs_time = obs_time, keepdim = True) # (nodes, 1) - T = data.T.item() - tI = tI.round().long() - tR = tR.round().long() - # compose a history - with torch.no_grad(): - y_pred = torch.zeros_like(data.y) # (nodes, T+1) - y_pred.scatter_(dim = 1, index = torch.minimum(tI, data.T), src = torch.full_like(tI, 1)) - y_pred.scatter_(dim = 1, index = torch.minimum(tR, data.T), src = torch.full_like(tR, 2)) - y_pred = y_pred[:, : data.T.item()].cummax(dim = 1).values - return y_pred - -args = get_args() -tester = Tester(args.data_dir, args.device, main) -tester.test([args.dataset], seed = args.seed, rep = 1) -tester.save(args.output) \ No newline at end of file diff --git a/experiments/exp3_beta_scatter.py b/experiments/exp3_beta_scatter.py new file mode 100644 index 0000000..061d875 --- /dev/null +++ b/experiments/exp3_beta_scatter.py @@ -0,0 +1,312 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +Experiment 3: Parameter Estimation Accuracy (beta scatter) + +Goal: + Draw ONE scatter plot: x = true beta, y = estimated beta_hat (from HERMES b_estim). + +Setting (simple & consistent with exp1/exp2): + - Synthetic diffusion on a fixed base graph (default: ba-si graph structure). + - Two observed snapshots for estimation: obs_time = [floor(T/2), T]. + - For each trial, sample a true beta ~ Uniform[beta_min, beta_max], simulate SI diffusion, + then estimate beta_hat using inc.diffus.b_estim (segmented mean-field pseudo-likelihood). + +Baseline (NOT plotted): + A very simple well-mixed/logistic estimator using infected fraction i_t: + beta_base = (logit(i_T) - logit(i_mid)) / (T - mid) + We only print its error metrics to show our estimator is more accurate. + +Run from repo root (example): + python experiments/exp3_beta_scatter.py --dataset ba-si --data_dir input --device cuda \ + --trials 50 --T 10 --beta_min 0.02 --beta_max 0.25 --b_steps 300 + +Outputs: + - CSV: output/exp3_beta_scatter/points.csv + - Plot: output/exp3_beta_scatter/beta_scatter.png (+ optional pdf) +""" + +from __future__ import annotations + +import os +import sys +import math +import csv +import argparse +from typing import List, Dict, Any, Tuple + +import numpy as np +import torch +from torch_geometric.data import Data + +# ------------------------------- +# Make repo root importable +# ------------------------------- +ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) +if ROOT not in sys.path: + sys.path.insert(0, ROOT) + +# Reuse existing project code (HERMES parameter estimation) +from inc.diffus import diffus_gen, b_estim, SIR_STATES # type: ignore +from inc.utils import seed_all # type: ignore + +# Plot (keep minimal, matplotlib only) +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt + + +def _ensure_dir(path: str) -> None: + os.makedirs(path, exist_ok=True) + + +def _logit(x: float, eps: float = 1e-6) -> float: + x = min(max(x, eps), 1.0 - eps) + return math.log(x / (1.0 - x)) + + +def beta_baseline_well_mixed(y: torch.Tensor, t1: int, t2: int) -> float: + """ + Simple baseline estimator using only infected fraction (no graph): + beta = (logit(i_t2) - logit(i_t1)) / (t2 - t1) + where i_t is fraction of infected nodes at time t. + """ + assert 0 <= t1 < t2 < y.size(1), "invalid t1/t2 for y shape (nodes, T+1)" + i1 = float((y[:, t1] == SIR_STATES.I).float().mean().item()) + i2 = float((y[:, t2] == SIR_STATES.I).float().mean().item()) + return (_logit(i2) - _logit(i1)) / float(t2 - t1) + + +def make_si_trial_data( + edge_index: torch.Tensor, + n_nodes: int, + T: int, + I0: int, + beta_true: float, + device: torch.device, +) -> Data: + """ + Simulate ONE SI diffusion history and pack into torch_geometric.data.Data: + data.y: (nodes, T+1) with states {S=0, I=1} + data.T: scalar tensor + """ + # diffus_gen returns Y: (T+1, nodes, samples) + Y = diffus_gen( + T=T, + n_nodes=n_nodes, + edge_index=edge_index, + I0=I0, + n_samples=1, + pI=beta_true, + pR=0.0, # SI + ) + y = Y[:, :, 0].T.contiguous() # (nodes, T+1) + + data = Data( + edge_index=edge_index, + y=y, + T=torch.tensor(T, dtype=torch.long, device=device), + ) + data.num_nodes = n_nodes + return data + + +def save_points_csv(path: str, rows: List[Dict[str, Any]]) -> None: + _ensure_dir(os.path.dirname(path)) + fieldnames = ["trial", "beta_true", "beta_hat", "beta_base"] + with open(path, "w", newline="", encoding="utf-8") as f: + w = csv.DictWriter(f, fieldnames=fieldnames) + w.writeheader() + for r in rows: + w.writerow({k: r.get(k, "") for k in fieldnames}) + + +def plot_scatter(path_png: str, xs: np.ndarray, ys: np.ndarray, beta_min: float, beta_max: float, title: str) -> None: + _ensure_dir(os.path.dirname(path_png)) + + # diagonal range + lo = min(beta_min, float(xs.min()), float(ys.min())) + hi = max(beta_max, float(xs.max()), float(ys.max())) + pad = 0.02 * (hi - lo + 1e-12) + lo, hi = lo - pad, hi + pad + + plt.figure(figsize=(5.2, 5.0)) + plt.scatter(xs, ys, s=18, alpha=0.8) + plt.plot([lo, hi], [lo, hi], linestyle="--", linewidth=1.2) + plt.xlim(lo, hi) + plt.ylim(lo, hi) + plt.xlabel(r"True $\beta$") + plt.ylabel(r"Estimated $\hat{\beta}$") + plt.title(title) + plt.tight_layout() + plt.savefig(path_png, dpi=200) + plt.close() + + +def main() -> None: + p = argparse.ArgumentParser() + + # base graph source (we only need edge_index & n) + p.add_argument("--dataset", type=str, default="ba-si", + help="use an existing dataset to load the base graph structure (e.g., ba-si, er-si).") + p.add_argument("--data_dir", type=str, default="input", + help="dataset folder (used by data loader in your repo).") + + # simulation controls + p.add_argument("--T", type=int, default=10) + p.add_argument("--trials", type=int, default=50) + p.add_argument("--beta_min", type=float, default=0.02) + p.add_argument("--beta_max", type=float, default=0.25) + p.add_argument("--seed", type=int, default=123456789) + p.add_argument("--I0_frac", type=float, default=0.05, + help="initial infected fraction for each simulated history (SI).") + + # device + p.add_argument("--device", type=str, default="cuda") + + # HERMES parameter estimation hyperparams (reuse b_estim) + p.add_argument("--b_pI0", type=float, default=0.05, + help="initial beta guess in b_estim") + p.add_argument("--b_pR0", type=float, default=0.0, + help="ignored for SI (kept for compatibility)") + p.add_argument("--b_steps", type=int, default=300) + p.add_argument("--b_lr", type=float, default=0.01) + + # output + p.add_argument("--out_dir", type=str, default="output/exp3_beta_scatter") + p.add_argument("--save_pdf", action="store_true") + + args = p.parse_args() + + # device resolve + if args.device.startswith("cuda") and not torch.cuda.is_available(): + print("[warn] CUDA not available, fallback to CPU.") + device = torch.device("cpu") + else: + device = torch.device(args.device) + + _ensure_dir(args.out_dir) + + # ---- load base graph structure ---- + # To avoid forcing this script to depend on extra libs, we try to load a cached .pt. + # In your repo, inc.data.data_load() handles caching; here we import it lazily. + try: + from inc.data import data_load # type: ignore + base = data_load(args.dataset, args.data_dir, device) + except Exception as e: + raise RuntimeError( + "Failed to load base dataset graph. " + "Please ensure your repo provides inc.data.data_load and the dataset cache exists." + ) from e + + edge_index = base.edge_index + n_nodes = int(base.num_nodes) + + # initial infected count + I0 = max(1, int(round(args.I0_frac * n_nodes))) + + # obs times: two frames + T = int(args.T) + obs_mid = T // 2 + obs_time = [obs_mid, T] + + # build a minimal args namespace for b_estim() + # b_estim expects args.b_pI0/b_pR0/b_steps/b_lr. + b_args = argparse.Namespace( + b_pI0=float(args.b_pI0), + b_pR0=float(args.b_pR0), + b_steps=int(args.b_steps), + b_lr=float(args.b_lr), + # keep obs_time for compatibility if your b_estim reads it + obs_time=str(obs_mid), + device=device, + ) + + # ---- run trials ---- + rows: List[Dict[str, Any]] = [] + betas_true: List[float] = [] + betas_hat: List[float] = [] + betas_base: List[float] = [] + + for k in range(args.trials): + seed_all(args.seed ^ k) + + beta_true = float(np.random.uniform(args.beta_min, args.beta_max)) + data_k = make_si_trial_data( + edge_index=edge_index, + n_nodes=n_nodes, + T=T, + I0=I0, + beta_true=beta_true, + device=device, + ) + + # estimate beta_hat using HERMES b_estim (segmented mean-field) + bpar = b_estim(data_k, b_args, obs_time=obs_time) + beta_hat = float(bpar["pI"] if isinstance(bpar, dict) else bpar.pI) + + # baseline (not plotted) + beta_base = float(beta_baseline_well_mixed(data_k.y, obs_mid, T)) + + rows.append({ + "trial": k, + "beta_true": beta_true, + "beta_hat": beta_hat, + "beta_base": beta_base, + }) + betas_true.append(beta_true) + betas_hat.append(beta_hat) + betas_base.append(beta_base) + + if (k + 1) % max(1, args.trials // 10) == 0: + print(f"[trial {k+1:>3d}/{args.trials}] beta={beta_true:.4f} hat={beta_hat:.4f} base={beta_base:.4f}") + + # ---- save csv ---- + csv_path = os.path.join(args.out_dir, "points.csv") + save_points_csv(csv_path, rows) + print(f"[ok] saved points -> {csv_path}") + + # ---- metrics (print only) ---- + bt = np.array(betas_true, dtype=float) + bh = np.array(betas_hat, dtype=float) + bb = np.array(betas_base, dtype=float) + + rmse_hat = float(np.sqrt(np.mean((bh - bt) ** 2))) + mae_hat = float(np.mean(np.abs(bh - bt))) + rmse_base = float(np.sqrt(np.mean((bb - bt) ** 2))) + mae_base = float(np.mean(np.abs(bb - bt))) + + print(f"[metric] ours RMSE={rmse_hat:.6f} MAE={mae_hat:.6f}") + print(f"[metric] base RMSE={rmse_base:.6f} MAE={mae_base:.6f} (not plotted)") + + # ---- plot (ONE scatter only: true beta vs beta_hat) ---- + title = rf"$T={T}$, obs={obs_time}, trials={args.trials} (RMSE={rmse_hat:.4f})" + fig_png = os.path.join(args.out_dir, "beta_scatter.png") + plot_scatter(fig_png, bt, bh, args.beta_min, args.beta_max, title) + print(f"[ok] saved figure -> {fig_png}") + + if args.save_pdf: + # re-render once to pdf (keep consistent) + fig_pdf = os.path.join(args.out_dir, "beta_scatter.pdf") + # quick replot + lo = min(args.beta_min, float(bt.min()), float(bh.min())) + hi = max(args.beta_max, float(bt.max()), float(bh.max())) + pad = 0.02 * (hi - lo + 1e-12) + lo, hi = lo - pad, hi + pad + plt.figure(figsize=(5.2, 5.0)) + plt.scatter(bt, bh, s=18, alpha=0.8) + plt.plot([lo, hi], [lo, hi], linestyle="--", linewidth=1.2) + plt.xlim(lo, hi) + plt.ylim(lo, hi) + plt.xlabel(r"True $\beta$") + plt.ylabel(r"Estimated $\hat{\beta}$") + plt.title(title) + plt.tight_layout() + plt.savefig(fig_pdf) + plt.close() + print(f"[ok] saved figure -> {fig_pdf}") + + +if __name__ == "__main__": + main() From 3351b715715c6c5f0c32609039bd98294922b0fd Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Thu, 5 Feb 2026 22:54:59 -0600 Subject: [PATCH 12/19] update exp3 --- experiments/exp3_beta_scatter.py | 67 +++++++++++++++++++++++--------- 1 file changed, 48 insertions(+), 19 deletions(-) diff --git a/experiments/exp3_beta_scatter.py b/experiments/exp3_beta_scatter.py index 061d875..019f1ee 100644 --- a/experiments/exp3_beta_scatter.py +++ b/experiments/exp3_beta_scatter.py @@ -20,7 +20,7 @@ Run from repo root (example): python experiments/exp3_beta_scatter.py --dataset ba-si --data_dir input --device cuda \ - --trials 50 --T 10 --beta_min 0.02 --beta_max 0.25 --b_steps 300 + --trials 300 --T 10 --beta_min 0.02 --beta_max 0.25 --b_steps 300 Outputs: - CSV: output/exp3_beta_scatter/points.csv @@ -122,26 +122,55 @@ def save_points_csv(path: str, rows: List[Dict[str, Any]]) -> None: w.writerow({k: r.get(k, "") for k in fieldnames}) -def plot_scatter(path_png: str, xs: np.ndarray, ys: np.ndarray, beta_min: float, beta_max: float, title: str) -> None: - _ensure_dir(os.path.dirname(path_png)) +import matplotlib as mpl +import matplotlib.pyplot as plt +import numpy as np + +def plot_scatter(beta_true, beta_hat, T, obs_time, out_png, rmse=None): + + mpl.rcParams.update({ + "font.size": 11, + "axes.labelsize": 13, + "xtick.labelsize": 11, + "ytick.labelsize": 11, + "axes.linewidth": 1.0, + }) - # diagonal range - lo = min(beta_min, float(xs.min()), float(ys.min())) - hi = max(beta_max, float(xs.max()), float(ys.max())) + beta_true = np.asarray(beta_true) + beta_hat = np.asarray(beta_hat) + + lo = float(min(beta_true.min(), beta_hat.min())) + hi = float(max(beta_true.max(), beta_hat.max())) pad = 0.02 * (hi - lo + 1e-12) lo, hi = lo - pad, hi + pad - plt.figure(figsize=(5.2, 5.0)) - plt.scatter(xs, ys, s=18, alpha=0.8) - plt.plot([lo, hi], [lo, hi], linestyle="--", linewidth=1.2) - plt.xlim(lo, hi) - plt.ylim(lo, hi) - plt.xlabel(r"True $\beta$") - plt.ylabel(r"Estimated $\hat{\beta}$") - plt.title(title) - plt.tight_layout() - plt.savefig(path_png, dpi=200) - plt.close() + fig, ax = plt.subplots(figsize=(3.2, 3.2), dpi=300) + + ax.scatter(beta_true, beta_hat, + s=10, alpha=0.65, linewidths=0) # s/alpha 你可微调 + + + ax.plot([lo, hi], [lo, hi], linestyle="--", color="0.4", linewidth=1.2) + + ax.set_xlabel(r"$\beta$") + ax.set_ylabel(r"$\hat{\beta}$") + ax.set_xlim(lo, hi) + ax.set_ylim(lo, hi) + + + ax.set_aspect("equal", adjustable="box") + + + info = fr"$T={T}$, obs={obs_time}" + if rmse is not None: + info += fr"\nRMSE={rmse:.4f}" + ax.text(0.05, 0.95, info, transform=ax.transAxes, + ha="left", va="top", fontsize=10) + + fig.tight_layout(pad=0.2) + fig.savefig(out_png, bbox_inches="tight") + plt.close(fig) + def main() -> None: @@ -170,8 +199,8 @@ def main() -> None: help="initial beta guess in b_estim") p.add_argument("--b_pR0", type=float, default=0.0, help="ignored for SI (kept for compatibility)") - p.add_argument("--b_steps", type=int, default=300) - p.add_argument("--b_lr", type=float, default=0.01) + p.add_argument("--b_steps", type=int, default=2000) + p.add_argument("--b_lr", type=float, default=0.001) # output p.add_argument("--out_dir", type=str, default="output/exp3_beta_scatter") From ae231c3c32327e874e1913d164a05573b657cd03 Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Sat, 7 Feb 2026 12:26:38 -0600 Subject: [PATCH 13/19] update exp3_0207 --- experiments/exp3_beta_scatter.py | 478 +++++++++++++++++-------------- 1 file changed, 255 insertions(+), 223 deletions(-) diff --git a/experiments/exp3_beta_scatter.py b/experiments/exp3_beta_scatter.py index 019f1ee..34a0eac 100644 --- a/experiments/exp3_beta_scatter.py +++ b/experiments/exp3_beta_scatter.py @@ -1,111 +1,134 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- -""" -Experiment 3: Parameter Estimation Accuracy (beta scatter) - -Goal: - Draw ONE scatter plot: x = true beta, y = estimated beta_hat (from HERMES b_estim). - -Setting (simple & consistent with exp1/exp2): - - Synthetic diffusion on a fixed base graph (default: ba-si graph structure). - - Two observed snapshots for estimation: obs_time = [floor(T/2), T]. - - For each trial, sample a true beta ~ Uniform[beta_min, beta_max], simulate SI diffusion, - then estimate beta_hat using inc.diffus.b_estim (segmented mean-field pseudo-likelihood). - -Baseline (NOT plotted): - A very simple well-mixed/logistic estimator using infected fraction i_t: - beta_base = (logit(i_T) - logit(i_mid)) / (T - mid) - We only print its error metrics to show our estimator is more accurate. - -Run from repo root (example): - python experiments/exp3_beta_scatter.py --dataset ba-si --data_dir input --device cuda \ - --trials 300 --T 10 --beta_min 0.02 --beta_max 0.25 --b_steps 300 - -Outputs: - - CSV: output/exp3_beta_scatter/points.csv - - Plot: output/exp3_beta_scatter/beta_scatter.png (+ optional pdf) -""" - -from __future__ import annotations import os import sys -import math import csv import argparse -from typing import List, Dict, Any, Tuple +import zlib +from typing import Dict, Any, List, Tuple, Optional import numpy as np import torch from torch_geometric.data import Data # ------------------------------- -# Make repo root importable +# Make repo root importable. +# This supports either layout: +# - /experiments/exp3_beta_scatter.py (inc/ in parent) +# - /exp3_beta_scatter.py (inc/ in same dir) # ------------------------------- -ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) -if ROOT not in sys.path: - sys.path.insert(0, ROOT) - -# Reuse existing project code (HERMES parameter estimation) +HERE = os.path.abspath(os.path.dirname(__file__)) +CANDIDATE_ROOTS = [HERE, os.path.abspath(os.path.join(HERE, ".."))] +for cand in CANDIDATE_ROOTS: + if os.path.isdir(os.path.join(cand, "inc")): + if cand not in sys.path: + sys.path.insert(0, cand) + break + +from inc.data import data_load # type: ignore from inc.diffus import diffus_gen, b_estim, SIR_STATES # type: ignore from inc.utils import seed_all # type: ignore -# Plot (keep minimal, matplotlib only) import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt +# ------------------------------- +# True beta settings for main datasets (D1/D2) +# ------------------------------- +# NOTE: +# True parameters are not stored in the cached .pt files, so we encode the *main experiment* +# settings here (DITTO paper Appendix D.1; our data.py follows the same settings). +def true_params_for_dataset(dataset: str) -> Tuple[float, float]: + d = dataset.lower() + is_sir = d.endswith("-sir") + is_si = d.endswith("-si") + if not (is_sir or is_si): + raise ValueError("dataset name must end with -si or -sir, got: %s" % dataset) + + # D1: synthetic graphs (BA/ER) + if d.startswith("ba-") or d.startswith("er-"): + pI = 0.1 + pR = 0.1 if is_sir else 0.0 + return pI, pR + + # D2: real graphs (Oregon2/Prost) with synthetic diffusion + if d.startswith("oregon2-") or d.startswith("prost-"): + pI = 0.1 + pR = 0.05 if is_sir else 0.0 + return pI, pR + + raise ValueError( + "Unsupported dataset for this exp3 scatter (needs synthetic diffusion with known beta): %s" + % dataset + ) + + def _ensure_dir(path: str) -> None: os.makedirs(path, exist_ok=True) -def _logit(x: float, eps: float = 1e-6) -> float: - x = min(max(x, eps), 1.0 - eps) - return math.log(x / (1.0 - x)) +def _stable_int_hash(s: str) -> int: + """Stable 32-bit hash (avoid Python's randomized hash).""" + return int(zlib.crc32(s.encode("utf-8")) & 0xFFFFFFFF) -def beta_baseline_well_mixed(y: torch.Tensor, t1: int, t2: int) -> float: - """ - Simple baseline estimator using only infected fraction (no graph): - beta = (logit(i_t2) - logit(i_t1)) / (t2 - t1) - where i_t is fraction of infected nodes at time t. - """ - assert 0 <= t1 < t2 < y.size(1), "invalid t1/t2 for y shape (nodes, T+1)" - i1 = float((y[:, t1] == SIR_STATES.I).float().mean().item()) - i2 = float((y[:, t2] == SIR_STATES.I).float().mean().item()) - return (_logit(i2) - _logit(i1)) / float(t2 - t1) - - -def make_si_trial_data( +def simulate_one_history( edge_index: torch.Tensor, n_nodes: int, T: int, I0: int, - beta_true: float, + pI_true: float, + pR_true: float, device: torch.device, + max_tries: int = 50, ) -> Data: """ - Simulate ONE SI diffusion history and pack into torch_geometric.data.Data: - data.y: (nodes, T+1) with states {S=0, I=1} - data.T: scalar tensor + Simulate ONE diffusion history with fixed parameters (pI_true, pR_true). + + For SIR (pR_true > 0), we optionally resample until at least one node reaches R + (otherwise b_estim will infer n_cls=2 and skip estimating pR). + + Returns a torch_geometric Data with: + data.edge_index + data.y : (nodes, T+1) + data.T : scalar tensor """ - # diffus_gen returns Y: (T+1, nodes, samples) - Y = diffus_gen( - T=T, - n_nodes=n_nodes, - edge_index=edge_index, - I0=I0, - n_samples=1, - pI=beta_true, - pR=0.0, # SI - ) - y = Y[:, :, 0].T.contiguous() # (nodes, T+1) + last_y: Optional[torch.Tensor] = None + + for _try in range(max_tries): + Y = diffus_gen( + T=T, + n_nodes=n_nodes, + edge_index=edge_index, + I0=I0, + n_samples=1, + pI=pI_true, + pR=pR_true, + ) # (T+1, nodes, 1) + + y = Y[:, :, 0].T.contiguous() # (nodes, T+1) + last_y = y + if pR_true > 0: + # Ensure recovered exists at final time so b_estim treats it as SIR. + if not (y == SIR_STATES.R).any().item(): + continue + + data = Data( + edge_index=edge_index, + y=y, + T=torch.tensor(T, dtype=torch.long, device=device), + ) + data.num_nodes = n_nodes + return data + + # Fallback: return the last simulated history even if it has no recovered state. + assert last_y is not None data = Data( edge_index=edge_index, - y=y, + y=last_y, T=torch.tensor(T, dtype=torch.long, device=device), ) data.num_nodes = n_nodes @@ -114,7 +137,7 @@ def make_si_trial_data( def save_points_csv(path: str, rows: List[Dict[str, Any]]) -> None: _ensure_dir(os.path.dirname(path)) - fieldnames = ["trial", "beta_true", "beta_hat", "beta_base"] + fieldnames = ["dataset", "trial", "param", "beta_true", "beta_hat"] with open(path, "w", newline="", encoding="utf-8") as f: w = csv.DictWriter(f, fieldnames=fieldnames) w.writeheader() @@ -122,219 +145,228 @@ def save_points_csv(path: str, rows: List[Dict[str, Any]]) -> None: w.writerow({k: r.get(k, "") for k in fieldnames}) -import matplotlib as mpl -import matplotlib.pyplot as plt -import numpy as np - -def plot_scatter(beta_true, beta_hat, T, obs_time, out_png, rmse=None): +def plot_scatter( + rows: List[Dict[str, Any]], + out_png: str, + title: str, + rmse: Optional[float] = None, + save_pdf: bool = False, +) -> None: + beta_true = np.asarray([float(r["beta_true"]) for r in rows], dtype=float) + beta_hat = np.asarray([float(r["beta_hat"]) for r in rows], dtype=float) + params = [str(r["param"]) for r in rows] - mpl.rcParams.update({ - "font.size": 11, - "axes.labelsize": 13, - "xtick.labelsize": 11, - "ytick.labelsize": 11, - "axes.linewidth": 1.0, - }) - - beta_true = np.asarray(beta_true) - beta_hat = np.asarray(beta_hat) + mask_I = np.array([p == "pI" for p in params], dtype=bool) + mask_R = np.array([p == "pR" for p in params], dtype=bool) lo = float(min(beta_true.min(), beta_hat.min())) hi = float(max(beta_true.max(), beta_hat.max())) - pad = 0.02 * (hi - lo + 1e-12) + pad = 0.05 * (hi - lo + 1e-12) lo, hi = lo - pad, hi + pad - fig, ax = plt.subplots(figsize=(3.2, 3.2), dpi=300) - - ax.scatter(beta_true, beta_hat, - s=10, alpha=0.65, linewidths=0) # s/alpha 你可微调 + fig, ax = plt.subplots(figsize=(3.4, 3.4), dpi=300) + if mask_I.any(): + ax.scatter(beta_true[mask_I], beta_hat[mask_I], s=10, alpha=0.65, linewidths=0, label=r"$\beta_I$") + if mask_R.any(): + ax.scatter(beta_true[mask_R], beta_hat[mask_R], s=14, alpha=0.65, linewidths=0, marker="x", label=r"$\beta_R$") ax.plot([lo, hi], [lo, hi], linestyle="--", color="0.4", linewidth=1.2) - ax.set_xlabel(r"$\beta$") - ax.set_ylabel(r"$\hat{\beta}$") + ax.set_xlabel(r"True $\beta$") + ax.set_ylabel(r"Estimated $\hat{\beta}$") ax.set_xlim(lo, hi) ax.set_ylim(lo, hi) - - ax.set_aspect("equal", adjustable="box") - - info = fr"$T={T}$, obs={obs_time}" + info = title if rmse is not None: - info += fr"\nRMSE={rmse:.4f}" - ax.text(0.05, 0.95, info, transform=ax.transAxes, - ha="left", va="top", fontsize=10) + info += "\nRMSE=%.4f" % rmse + ax.text(0.05, 0.95, info, transform=ax.transAxes, ha="left", va="top", fontsize=9) + + if mask_R.any(): + ax.legend(frameon=False, fontsize=9, loc="lower right") fig.tight_layout(pad=0.2) fig.savefig(out_png, bbox_inches="tight") plt.close(fig) + if save_pdf: + out_pdf = os.path.splitext(out_png)[0] + ".pdf" + fig, ax = plt.subplots(figsize=(3.4, 3.4), dpi=300) + if mask_I.any(): + ax.scatter(beta_true[mask_I], beta_hat[mask_I], s=10, alpha=0.65, linewidths=0, label=r"$\beta_I$") + if mask_R.any(): + ax.scatter(beta_true[mask_R], beta_hat[mask_R], s=14, alpha=0.65, linewidths=0, marker="x", label=r"$\beta_R$") + ax.plot([lo, hi], [lo, hi], linestyle="--", color="0.4", linewidth=1.2) + ax.set_xlabel(r"True $\beta$") + ax.set_ylabel(r"Estimated $\hat{\beta}$") + ax.set_xlim(lo, hi) + ax.set_ylim(lo, hi) + ax.set_aspect("equal", adjustable="box") + if mask_R.any(): + ax.legend(frameon=False, fontsize=9, loc="lower right") + fig.tight_layout(pad=0.2) + fig.savefig(out_pdf, bbox_inches="tight") + plt.close(fig) + + +def parse_datasets_arg(s: str) -> List[str]: + s = (s or "").strip() + if s.lower() in {"main", "default"}: + # 8 synthetic datasets in main experiments (D1+D2), where true beta is known. + return [ + "ba-si", "ba-sir", + "er-si", "er-sir", + "oregon2-si", "oregon2-sir", + "prost-si", "prost-sir", + ] + return [x.strip() for x in s.split(",") if x.strip()] def main() -> None: p = argparse.ArgumentParser() - # base graph source (we only need edge_index & n) - p.add_argument("--dataset", type=str, default="ba-si", - help="use an existing dataset to load the base graph structure (e.g., ba-si, er-si).") - p.add_argument("--data_dir", type=str, default="input", - help="dataset folder (used by data loader in your repo).") - - # simulation controls - p.add_argument("--T", type=int, default=10) - p.add_argument("--trials", type=int, default=50) - p.add_argument("--beta_min", type=float, default=0.02) - p.add_argument("--beta_max", type=float, default=0.25) + p.add_argument( + "--datasets", + type=str, + default="main", + help="Comma-separated datasets, or 'main' for BA/ER/Oregon2/Prost with SI+SIR.", + ) + p.add_argument("--data_dir", type=str, default="input") + p.add_argument("--device", type=str, default="cuda") + + # trials + p.add_argument("--trials_per_dataset", type=int, default=30) p.add_argument("--seed", type=int, default=123456789) - p.add_argument("--I0_frac", type=float, default=0.05, - help="initial infected fraction for each simulated history (SI).") - # device - p.add_argument("--device", type=str, default="cuda") + # plotting: default includes beta_R for SIR + p.add_argument( + "--only_betaI", + action="store_true", + help="If set, only plot infection beta_I (skip recovery beta_R).", + ) - # HERMES parameter estimation hyperparams (reuse b_estim) - p.add_argument("--b_pI0", type=float, default=0.05, - help="initial beta guess in b_estim") - p.add_argument("--b_pR0", type=float, default=0.0, - help="ignored for SI (kept for compatibility)") - p.add_argument("--b_steps", type=int, default=2000) + # b_estim hyperparams + p.add_argument("--b_pI0", type=float, default=0.05) + p.add_argument("--b_pR0", type=float, default=0.05) + p.add_argument("--b_steps", type=int, default=300) p.add_argument("--b_lr", type=float, default=0.001) # output - p.add_argument("--out_dir", type=str, default="output/exp3_beta_scatter") + p.add_argument("--out_dir", type=str, default="output/exp3_beta_scatter_main") p.add_argument("--save_pdf", action="store_true") args = p.parse_args() # device resolve - if args.device.startswith("cuda") and not torch.cuda.is_available(): + if args.device.startswith("cuda") and (not torch.cuda.is_available()): print("[warn] CUDA not available, fallback to CPU.") device = torch.device("cpu") else: device = torch.device(args.device) + datasets = parse_datasets_arg(args.datasets) + if not datasets: + raise ValueError("Empty --datasets") + _ensure_dir(args.out_dir) - # ---- load base graph structure ---- - # To avoid forcing this script to depend on extra libs, we try to load a cached .pt. - # In your repo, inc.data.data_load() handles caching; here we import it lazily. - try: - from inc.data import data_load # type: ignore - base = data_load(args.dataset, args.data_dir, device) - except Exception as e: - raise RuntimeError( - "Failed to load base dataset graph. " - "Please ensure your repo provides inc.data.data_load and the dataset cache exists." - ) from e - - edge_index = base.edge_index - n_nodes = int(base.num_nodes) - - # initial infected count - I0 = max(1, int(round(args.I0_frac * n_nodes))) - - # obs times: two frames - T = int(args.T) - obs_mid = T // 2 - obs_time = [obs_mid, T] - - # build a minimal args namespace for b_estim() - # b_estim expects args.b_pI0/b_pR0/b_steps/b_lr. + # b_estim expects these in args; we also pass obs_time explicitly. b_args = argparse.Namespace( b_pI0=float(args.b_pI0), b_pR0=float(args.b_pR0), b_steps=int(args.b_steps), b_lr=float(args.b_lr), - # keep obs_time for compatibility if your b_estim reads it - obs_time=str(obs_mid), + obs_time="", # not used when obs_time is passed explicitly device=device, ) - # ---- run trials ---- rows: List[Dict[str, Any]] = [] - betas_true: List[float] = [] - betas_hat: List[float] = [] - betas_base: List[float] = [] - for k in range(args.trials): - seed_all(args.seed ^ k) - - beta_true = float(np.random.uniform(args.beta_min, args.beta_max)) - data_k = make_si_trial_data( - edge_index=edge_index, - n_nodes=n_nodes, - T=T, - I0=I0, - beta_true=beta_true, - device=device, + for dataset in datasets: + # Load once to get base graph (edge_index, num_nodes) and T. + base = data_load(dataset, args.data_dir, device) + edge_index = base.edge_index + n_nodes = int(base.num_nodes) + T = int(base.T.item()) + + # main protocol: observe { floor(T/2), T } + t_obs = T // 2 + obs_time = [t_obs, T] + + # I0 from cached dataset (matches generation protocol) + I0 = int((base.y[:, 0] == SIR_STATES.I).sum().item()) + I0 = max(1, I0) + + pI_true, pR_true = true_params_for_dataset(dataset) + + # run multiple trials on this dataset + for k in range(int(args.trials_per_dataset)): + seed_k = (int(args.seed) ^ _stable_int_hash(dataset) ^ int(k)) & 0xFFFFFFFF + seed_all(seed_k) + + data_k = simulate_one_history( + edge_index=edge_index, + n_nodes=n_nodes, + T=T, + I0=I0, + pI_true=pI_true, + pR_true=pR_true, + device=device, + ) + + est = b_estim(data_k, b_args, obs_time=obs_time) # dict with keys pI, pR + + # record beta_I + rows.append( + dict( + dataset=dataset, + trial=k, + param="pI", + beta_true=float(pI_true), + beta_hat=float(est["pI"]), + ) + ) + + # record beta_R for SIR unless disabled + if (not args.only_betaI) and (pR_true > 0): + rows.append( + dict( + dataset=dataset, + trial=k, + param="pR", + beta_true=float(pR_true), + beta_hat=float(est.get("pR", 0.0)), + ) + ) + + print( + "[ok] %s: n=%d, T=%d, obs=%s, I0=%d, true=(pI=%.3f, pR=%.3f), trials=%d" + % (dataset, n_nodes, T, str(obs_time), I0, pI_true, pR_true, int(args.trials_per_dataset)) ) - # estimate beta_hat using HERMES b_estim (segmented mean-field) - bpar = b_estim(data_k, b_args, obs_time=obs_time) - beta_hat = float(bpar["pI"] if isinstance(bpar, dict) else bpar.pI) - - # baseline (not plotted) - beta_base = float(beta_baseline_well_mixed(data_k.y, obs_mid, T)) - - rows.append({ - "trial": k, - "beta_true": beta_true, - "beta_hat": beta_hat, - "beta_base": beta_base, - }) - betas_true.append(beta_true) - betas_hat.append(beta_hat) - betas_base.append(beta_base) - - if (k + 1) % max(1, args.trials // 10) == 0: - print(f"[trial {k+1:>3d}/{args.trials}] beta={beta_true:.4f} hat={beta_hat:.4f} base={beta_base:.4f}") - # ---- save csv ---- csv_path = os.path.join(args.out_dir, "points.csv") save_points_csv(csv_path, rows) - print(f"[ok] saved points -> {csv_path}") + print("[ok] saved points -> %s" % csv_path) - # ---- metrics (print only) ---- - bt = np.array(betas_true, dtype=float) - bh = np.array(betas_hat, dtype=float) - bb = np.array(betas_base, dtype=float) + # ---- metrics ---- + bt = np.array([float(r["beta_true"]) for r in rows], dtype=float) + bh = np.array([float(r["beta_hat"]) for r in rows], dtype=float) + rmse = float(np.sqrt(np.mean((bh - bt) ** 2))) + mae = float(np.mean(np.abs(bh - bt))) + print("[metric] overall RMSE=%.6f MAE=%.6f (points=%d)" % (rmse, mae, len(rows))) - rmse_hat = float(np.sqrt(np.mean((bh - bt) ** 2))) - mae_hat = float(np.mean(np.abs(bh - bt))) - rmse_base = float(np.sqrt(np.mean((bb - bt) ** 2))) - mae_base = float(np.mean(np.abs(bb - bt))) - - print(f"[metric] ours RMSE={rmse_hat:.6f} MAE={mae_hat:.6f}") - print(f"[metric] base RMSE={rmse_base:.6f} MAE={mae_base:.6f} (not plotted)") - - # ---- plot (ONE scatter only: true beta vs beta_hat) ---- - title = rf"$T={T}$, obs={obs_time}, trials={args.trials} (RMSE={rmse_hat:.4f})" + # ---- plot ---- fig_png = os.path.join(args.out_dir, "beta_scatter.png") - plot_scatter(fig_png, bt, bh, args.beta_min, args.beta_max, title) - print(f"[ok] saved figure -> {fig_png}") - - if args.save_pdf: - # re-render once to pdf (keep consistent) - fig_pdf = os.path.join(args.out_dir, "beta_scatter.pdf") - # quick replot - lo = min(args.beta_min, float(bt.min()), float(bh.min())) - hi = max(args.beta_max, float(bt.max()), float(bh.max())) - pad = 0.02 * (hi - lo + 1e-12) - lo, hi = lo - pad, hi + pad - plt.figure(figsize=(5.2, 5.0)) - plt.scatter(bt, bh, s=18, alpha=0.8) - plt.plot([lo, hi], [lo, hi], linestyle="--", linewidth=1.2) - plt.xlim(lo, hi) - plt.ylim(lo, hi) - plt.xlabel(r"True $\beta$") - plt.ylabel(r"Estimated $\hat{\beta}$") - plt.title(title) - plt.tight_layout() - plt.savefig(fig_pdf) - plt.close() - print(f"[ok] saved figure -> {fig_pdf}") + title = "datasets=%d, trials/ds=%d, only_betaI=%s" % ( + len(datasets), + int(args.trials_per_dataset), + str(bool(args.only_betaI)), + ) + plot_scatter(rows, out_png=fig_png, title=title, rmse=rmse, save_pdf=bool(args.save_pdf)) + print("[ok] saved figure -> %s" % fig_png) if __name__ == "__main__": From c1e5abe5989772effdba89be9eb62f94f7688bf4 Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Sat, 7 Feb 2026 14:31:47 -0600 Subject: [PATCH 14/19] update exp3 csv only --- experiments/exp3_beta_scatter.py | 339 ++++++++----------------------- 1 file changed, 89 insertions(+), 250 deletions(-) diff --git a/experiments/exp3_beta_scatter.py b/experiments/exp3_beta_scatter.py index 34a0eac..4948668 100644 --- a/experiments/exp3_beta_scatter.py +++ b/experiments/exp3_beta_scatter.py @@ -1,22 +1,11 @@ - - import os import sys import csv import argparse -import zlib -from typing import Dict, Any, List, Tuple, Optional +from typing import Dict, Any, List, Tuple -import numpy as np import torch -from torch_geometric.data import Data -# ------------------------------- -# Make repo root importable. -# This supports either layout: -# - /experiments/exp3_beta_scatter.py (inc/ in parent) -# - /exp3_beta_scatter.py (inc/ in same dir) -# ------------------------------- HERE = os.path.abspath(os.path.dirname(__file__)) CANDIDATE_ROOTS = [HERE, os.path.abspath(os.path.join(HERE, ".."))] for cand in CANDIDATE_ROOTS: @@ -26,26 +15,54 @@ break from inc.data import data_load # type: ignore -from inc.diffus import diffus_gen, b_estim, SIR_STATES # type: ignore -from inc.utils import seed_all # type: ignore +from inc.diffus import b_estim # type: ignore -import matplotlib -matplotlib.use("Agg") -import matplotlib.pyplot as plt + +def _ensure_dir(path: str) -> None: + os.makedirs(path, exist_ok=True) + + +def parse_datasets_arg(s: str) -> List[str]: + s = (s or "").strip() + if s.lower() in {"main", "default"}: + # 8 datasets used in main experiments where true parameters are known by construction. + return [ + "ba-si", "ba-sir", + "er-si", "er-sir", + "oregon2-si", "oregon2-sir", + "prost-si", "prost-sir", + ] + return [x.strip() for x in s.split(",") if x.strip()] + + +def parse_obs_time_arg(s: str, T: int) -> List[int]: + """ + Parse --obs_time. If empty, use main protocol {floor(T/2), T}. + Input format examples: + --obs_time "" -> [T//2, T] + --obs_time "5" -> [5, T] + --obs_time "5,10" -> [5, 10] (and ensures T included) + """ + s = (s or "").strip() + if not s: + return [T // 2, T] + ts = [int(x) for x in s.split(",") if x.strip() != ""] + ts.append(T) + ts = sorted(set(t for t in ts if 0 <= t <= T)) + if ts[-1] != T: + ts.append(T) + return ts # ------------------------------- -# True beta settings for main datasets (D1/D2) +# Justified true parameters for MAIN synthetic datasets # ------------------------------- -# NOTE: -# True parameters are not stored in the cached .pt files, so we encode the *main experiment* -# settings here (DITTO paper Appendix D.1; our data.py follows the same settings). def true_params_for_dataset(dataset: str) -> Tuple[float, float]: d = dataset.lower() is_sir = d.endswith("-sir") is_si = d.endswith("-si") if not (is_sir or is_si): - raise ValueError("dataset name must end with -si or -sir, got: %s" % dataset) + raise ValueError(f"dataset must end with -si or -sir, got: {dataset}") # D1: synthetic graphs (BA/ER) if d.startswith("ba-") or d.startswith("er-"): @@ -60,84 +77,13 @@ def true_params_for_dataset(dataset: str) -> Tuple[float, float]: return pI, pR raise ValueError( - "Unsupported dataset for this exp3 scatter (needs synthetic diffusion with known beta): %s" - % dataset + f"Unsupported dataset for exp3 load-only scatter (needs known true beta/gamma): {dataset}" ) -def _ensure_dir(path: str) -> None: - os.makedirs(path, exist_ok=True) - - -def _stable_int_hash(s: str) -> int: - """Stable 32-bit hash (avoid Python's randomized hash).""" - return int(zlib.crc32(s.encode("utf-8")) & 0xFFFFFFFF) - - -def simulate_one_history( - edge_index: torch.Tensor, - n_nodes: int, - T: int, - I0: int, - pI_true: float, - pR_true: float, - device: torch.device, - max_tries: int = 50, -) -> Data: - """ - Simulate ONE diffusion history with fixed parameters (pI_true, pR_true). - - For SIR (pR_true > 0), we optionally resample until at least one node reaches R - (otherwise b_estim will infer n_cls=2 and skip estimating pR). - - Returns a torch_geometric Data with: - data.edge_index - data.y : (nodes, T+1) - data.T : scalar tensor - """ - last_y: Optional[torch.Tensor] = None - - for _try in range(max_tries): - Y = diffus_gen( - T=T, - n_nodes=n_nodes, - edge_index=edge_index, - I0=I0, - n_samples=1, - pI=pI_true, - pR=pR_true, - ) # (T+1, nodes, 1) - - y = Y[:, :, 0].T.contiguous() # (nodes, T+1) - last_y = y - - if pR_true > 0: - # Ensure recovered exists at final time so b_estim treats it as SIR. - if not (y == SIR_STATES.R).any().item(): - continue - - data = Data( - edge_index=edge_index, - y=y, - T=torch.tensor(T, dtype=torch.long, device=device), - ) - data.num_nodes = n_nodes - return data - - # Fallback: return the last simulated history even if it has no recovered state. - assert last_y is not None - data = Data( - edge_index=edge_index, - y=last_y, - T=torch.tensor(T, dtype=torch.long, device=device), - ) - data.num_nodes = n_nodes - return data - - def save_points_csv(path: str, rows: List[Dict[str, Any]]) -> None: _ensure_dir(os.path.dirname(path)) - fieldnames = ["dataset", "trial", "param", "beta_true", "beta_hat"] + fieldnames = ["dataset", "param", "beta_true", "beta_hat", "T", "obs_time"] with open(path, "w", newline="", encoding="utf-8") as f: w = csv.DictWriter(f, fieldnames=fieldnames) w.writeheader() @@ -145,85 +91,6 @@ def save_points_csv(path: str, rows: List[Dict[str, Any]]) -> None: w.writerow({k: r.get(k, "") for k in fieldnames}) -def plot_scatter( - rows: List[Dict[str, Any]], - out_png: str, - title: str, - rmse: Optional[float] = None, - save_pdf: bool = False, -) -> None: - beta_true = np.asarray([float(r["beta_true"]) for r in rows], dtype=float) - beta_hat = np.asarray([float(r["beta_hat"]) for r in rows], dtype=float) - params = [str(r["param"]) for r in rows] - - mask_I = np.array([p == "pI" for p in params], dtype=bool) - mask_R = np.array([p == "pR" for p in params], dtype=bool) - - lo = float(min(beta_true.min(), beta_hat.min())) - hi = float(max(beta_true.max(), beta_hat.max())) - pad = 0.05 * (hi - lo + 1e-12) - lo, hi = lo - pad, hi + pad - - fig, ax = plt.subplots(figsize=(3.4, 3.4), dpi=300) - - if mask_I.any(): - ax.scatter(beta_true[mask_I], beta_hat[mask_I], s=10, alpha=0.65, linewidths=0, label=r"$\beta_I$") - if mask_R.any(): - ax.scatter(beta_true[mask_R], beta_hat[mask_R], s=14, alpha=0.65, linewidths=0, marker="x", label=r"$\beta_R$") - - ax.plot([lo, hi], [lo, hi], linestyle="--", color="0.4", linewidth=1.2) - - ax.set_xlabel(r"True $\beta$") - ax.set_ylabel(r"Estimated $\hat{\beta}$") - ax.set_xlim(lo, hi) - ax.set_ylim(lo, hi) - ax.set_aspect("equal", adjustable="box") - - info = title - if rmse is not None: - info += "\nRMSE=%.4f" % rmse - ax.text(0.05, 0.95, info, transform=ax.transAxes, ha="left", va="top", fontsize=9) - - if mask_R.any(): - ax.legend(frameon=False, fontsize=9, loc="lower right") - - fig.tight_layout(pad=0.2) - fig.savefig(out_png, bbox_inches="tight") - plt.close(fig) - - if save_pdf: - out_pdf = os.path.splitext(out_png)[0] + ".pdf" - fig, ax = plt.subplots(figsize=(3.4, 3.4), dpi=300) - if mask_I.any(): - ax.scatter(beta_true[mask_I], beta_hat[mask_I], s=10, alpha=0.65, linewidths=0, label=r"$\beta_I$") - if mask_R.any(): - ax.scatter(beta_true[mask_R], beta_hat[mask_R], s=14, alpha=0.65, linewidths=0, marker="x", label=r"$\beta_R$") - ax.plot([lo, hi], [lo, hi], linestyle="--", color="0.4", linewidth=1.2) - ax.set_xlabel(r"True $\beta$") - ax.set_ylabel(r"Estimated $\hat{\beta}$") - ax.set_xlim(lo, hi) - ax.set_ylim(lo, hi) - ax.set_aspect("equal", adjustable="box") - if mask_R.any(): - ax.legend(frameon=False, fontsize=9, loc="lower right") - fig.tight_layout(pad=0.2) - fig.savefig(out_pdf, bbox_inches="tight") - plt.close(fig) - - -def parse_datasets_arg(s: str) -> List[str]: - s = (s or "").strip() - if s.lower() in {"main", "default"}: - # 8 synthetic datasets in main experiments (D1+D2), where true beta is known. - return [ - "ba-si", "ba-sir", - "er-si", "er-sir", - "oregon2-si", "oregon2-sir", - "prost-si", "prost-sir", - ] - return [x.strip() for x in s.split(",") if x.strip()] - - def main() -> None: p = argparse.ArgumentParser() @@ -236,15 +103,12 @@ def main() -> None: p.add_argument("--data_dir", type=str, default="input") p.add_argument("--device", type=str, default="cuda") - # trials - p.add_argument("--trials_per_dataset", type=int, default=30) - p.add_argument("--seed", type=int, default=123456789) - - # plotting: default includes beta_R for SIR + # Observation times: if empty, use main protocol {floor(T/2), T} p.add_argument( - "--only_betaI", - action="store_true", - help="If set, only plot infection beta_I (skip recovery beta_R).", + "--obs_time", + type=str, + default="", + help='Observation frames used for estimation. "" means {floor(T/2), T}. Example: "5" or "5,10".', ) # b_estim hyperparams @@ -255,7 +119,14 @@ def main() -> None: # output p.add_argument("--out_dir", type=str, default="output/exp3_beta_scatter_main") - p.add_argument("--save_pdf", action="store_true") + p.add_argument("--csv_name", type=str, default="points.csv") + + # whether to also output pR (gamma) for SIR datasets + p.add_argument( + "--only_betaI", + action="store_true", + help="If set, only output infection beta (pI).", + ) args = p.parse_args() @@ -272,7 +143,7 @@ def main() -> None: _ensure_dir(args.out_dir) - # b_estim expects these in args; we also pass obs_time explicitly. + # b_estim expects these fields on args; we pass obs_time explicitly anyway. b_args = argparse.Namespace( b_pI0=float(args.b_pI0), b_pR0=float(args.b_pR0), @@ -285,88 +156,56 @@ def main() -> None: rows: List[Dict[str, Any]] = [] for dataset in datasets: - # Load once to get base graph (edge_index, num_nodes) and T. - base = data_load(dataset, args.data_dir, device) - edge_index = base.edge_index - n_nodes = int(base.num_nodes) - T = int(base.T.item()) - - # main protocol: observe { floor(T/2), T } - t_obs = T // 2 - obs_time = [t_obs, T] - - # I0 from cached dataset (matches generation protocol) - I0 = int((base.y[:, 0] == SIR_STATES.I).sum().item()) - I0 = max(1, I0) + data = data_load(dataset, args.data_dir, device) + T = int(data.T.item()) + obs_time = parse_obs_time_arg(args.obs_time, T) + obs_time_str = ",".join(str(t) for t in obs_time) pI_true, pR_true = true_params_for_dataset(dataset) - # run multiple trials on this dataset - for k in range(int(args.trials_per_dataset)): - seed_k = (int(args.seed) ^ _stable_int_hash(dataset) ^ int(k)) & 0xFFFFFFFF - seed_all(seed_k) + est = b_estim(data, b_args, obs_time=obs_time) # dict with keys pI, pR - data_k = simulate_one_history( - edge_index=edge_index, - n_nodes=n_nodes, + # always output infection beta (pI) + rows.append( + dict( + dataset=dataset, + param="pI", + beta_true=float(pI_true), + beta_hat=float(est.get("pI", float("nan"))), T=T, - I0=I0, - pI_true=pI_true, - pR_true=pR_true, - device=device, + obs_time=obs_time_str, ) + ) - est = b_estim(data_k, b_args, obs_time=obs_time) # dict with keys pI, pR - - # record beta_I + # optionally output recovery beta (pR) for SIR datasets + if (not args.only_betaI) and dataset.lower().endswith("-sir"): rows.append( dict( dataset=dataset, - trial=k, - param="pI", - beta_true=float(pI_true), - beta_hat=float(est["pI"]), + param="pR", + beta_true=float(pR_true), + beta_hat=float(est.get("pR", float("nan"))), + T=T, + obs_time=obs_time_str, ) ) - # record beta_R for SIR unless disabled - if (not args.only_betaI) and (pR_true > 0): - rows.append( - dict( - dataset=dataset, - trial=k, - param="pR", - beta_true=float(pR_true), - beta_hat=float(est.get("pR", 0.0)), - ) - ) - print( - "[ok] %s: n=%d, T=%d, obs=%s, I0=%d, true=(pI=%.3f, pR=%.3f), trials=%d" - % (dataset, n_nodes, T, str(obs_time), I0, pI_true, pR_true, int(args.trials_per_dataset)) + "[ok] %s: T=%d, obs={%s}, true(pI=%.3f, pR=%.3f) -> hat(pI=%.4f, pR=%.4f)" + % ( + dataset, + T, + obs_time_str, + pI_true, + pR_true, + float(est.get("pI", 0.0)), + float(est.get("pR", 0.0)), + ) ) - # ---- save csv ---- - csv_path = os.path.join(args.out_dir, "points.csv") - save_points_csv(csv_path, rows) - print("[ok] saved points -> %s" % csv_path) - - # ---- metrics ---- - bt = np.array([float(r["beta_true"]) for r in rows], dtype=float) - bh = np.array([float(r["beta_hat"]) for r in rows], dtype=float) - rmse = float(np.sqrt(np.mean((bh - bt) ** 2))) - mae = float(np.mean(np.abs(bh - bt))) - print("[metric] overall RMSE=%.6f MAE=%.6f (points=%d)" % (rmse, mae, len(rows))) - - # ---- plot ---- - fig_png = os.path.join(args.out_dir, "beta_scatter.png") - title = "datasets=%d, trials/ds=%d, only_betaI=%s" % ( - len(datasets), - int(args.trials_per_dataset), - str(bool(args.only_betaI)), - ) - plot_scatter(rows, out_png=fig_png, title=title, rmse=rmse, save_pdf=bool(args.save_pdf)) - print("[ok] saved figure -> %s" % fig_png) + out_csv = os.path.join(args.out_dir, args.csv_name) + save_points_csv(out_csv, rows) + print("[ok] saved csv -> %s (rows=%d)" % (out_csv, len(rows))) if __name__ == "__main__": From a367113e20f3f0ea2ca91b976c2ccb517417c86b Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Mon, 6 Apr 2026 16:15:14 -0500 Subject: [PATCH 15/19] Add HERMES nI0 sensitivity and MCMC diagnostics experiments --- experiments/exp_mcmc_diagnostics.py | 212 ++++++++++++++++++++++++++++ experiments/exp_nI0_sensitivity.py | 107 ++++++++++++++ hermes.py | 60 ++++++-- inc/diffus.py | 25 ++-- 4 files changed, 382 insertions(+), 22 deletions(-) create mode 100644 experiments/exp_mcmc_diagnostics.py create mode 100644 experiments/exp_nI0_sensitivity.py diff --git a/experiments/exp_mcmc_diagnostics.py b/experiments/exp_mcmc_diagnostics.py new file mode 100644 index 0000000..e64f43b --- /dev/null +++ b/experiments/exp_mcmc_diagnostics.py @@ -0,0 +1,212 @@ +import gc +import os +import sys +import os.path as osp +from copy import deepcopy + +import pandas as pd +import matplotlib.pyplot as plt + +ROOT = osp.dirname(osp.dirname(osp.abspath(__file__))) +if ROOT not in sys.path: + sys.path.insert(0, ROOT) + +from inc.data import * +from inc.test import TEST_METRICS +from hermes import b_estim, q_train, t_mcmc + +DEFAULT_T_STEPS = [25, 50, 100, 200] +DEFAULT_REPS = 4 + + +def get_args(): + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type = str, required = True, help = 'dataset name') + parser.add_argument('--seed', type = int, default = 12345, help = 'base seed used to fit bpar/q_net and derive MCMC seeds') + parser.add_argument('--data_dir', type = str, required = True, help = 'dataset folder') + parser.add_argument('--output_prefix', type = str, required = True, help = 'output prefix for csv/png files') + parser.add_argument('--device', type = torch.device, default = torch_device(), help = 'torch device') + + # same HERMES hyperparameters as hermes.py + parser.add_argument('--b_pI0', type = float, help = 'initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type = float, help = 'initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type = int, help = 'optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type = float, help = 'learning rate in diffusion parameter estimation') + parser.add_argument('--q_steps', type = int, help = 'training steps for the proposal model') + parser.add_argument('--q_lr', type = float, help = 'learning rate for the proposal model') + parser.add_argument('--q_hid', type = int, help = 'hidden size of the proposal model') + parser.add_argument('--q_gnn', type = int, help = 'number of layers of the GNN in the proposal model') + parser.add_argument('--q_mlp', type = int, help = 'number of layers of the MLP in the proposal model') + parser.add_argument('--q_samples', type = int, help = 'sample size to estimate the loss function of the proposal model') + parser.add_argument('--q_zlim', type = int, help = 'a hyperparameter to stablize gradient') + parser.add_argument('--p_coef', type = float, help = 'the coefficient gamma in the initial distribution P[y_0]') + parser.add_argument('--t_samples', type = int, help = 'MCMC sample size') + parser.add_argument('--t_steps', type = int, default = 100, help = 'unused default; sweep values are controlled internally') + parser.add_argument('--t_keep', type = float, help = 'moving average in MCMC') + parser.add_argument('--obs_time', type = str, default = '', help = 'extra observed snapshot times, comma-separated, e.g., 5,7,9') + return parser.parse_args() + + +def parse_obs_time(data, obs_time): + out = [int(t) for t in str(obs_time).split(',') if t] + out.append(data.T.item()) + return sorted(set(out)) + + +@torch.no_grad() +def compose_history(data, tI, tR): + tI = tI.round().long() + tR = tR.round().long() + y_pred = torch.zeros_like(data.y) + y_pred.scatter_(dim = 1, index = torch.minimum(tI, data.T), src = torch.full_like(tI, SIR_STATES.I)) + y_pred.scatter_(dim = 1, index = torch.minimum(tR, data.T), src = torch.full_like(tR, SIR_STATES.R)) + return y_pred[:, : data.T.item()].cummax(dim = 1).values + + +def save_trace_plot(df, y_col, ylabel, title, fpath): + plt.figure() + for t_steps, group in df.groupby('t_steps'): + mean_trace = group.groupby('step')[y_col].mean().reset_index() + plt.plot(mean_trace['step'], mean_trace[y_col], label = f'S={t_steps}') + plt.xlabel('MCMC step') + plt.ylabel(ylabel) + plt.title(title) + plt.legend() + plt.tight_layout() + plt.savefig(fpath, dpi = 200) + plt.close() + + +def save_metric_plot(df, metric, ylabel, title, fpath): + stats = df.groupby('t_steps')[metric].agg(['mean', 'std']).reset_index() + plt.figure() + plt.errorbar(stats['t_steps'], stats['mean'], yerr = stats['std'], marker = 'o') + plt.xlabel('t_steps (S)') + plt.ylabel(ylabel) + plt.title(title) + plt.tight_layout() + plt.savefig(fpath, dpi = 200) + plt.close() + + +def main(): + args = get_args() + out_dir = osp.dirname(args.output_prefix) + if out_dir: + os.makedirs(out_dir, exist_ok = True) + + data = data_load(args.dataset, args.data_dir, args.device) + obs_time = parse_obs_time(data, args.obs_time) + + # Fix parameter estimation + proposal training. + seed_all(args.seed) + bpar = b_estim(data, args, obs_time = obs_time) + print(f'[est] pI={bpar.pI:.4f}, pR={bpar.pR:.4f}', flush = True) + q_net = q_train(data, obs_time, bpar, args) + + trace_rows = [] + metric_rows = [] + mcmc_seeds = [args.seed ^ (rep + 1) for rep in range(DEFAULT_REPS)] + + for t_steps in DEFAULT_T_STEPS: + args_run = deepcopy(args) + args_run.t_steps = t_steps + for rep, mcmc_seed in enumerate(mcmc_seeds): + print(f'[dataset={args.dataset}] [S={t_steps}] [rep={rep}] [seed={mcmc_seed}]', flush = True) + seed_all(mcmc_seed) + tI, tR, diag_rows = t_mcmc( + data, + bpar, + q_net, + args_run, + obs_time = obs_time, + keepdim = True, + diagnostics = True, + ) + + y_pred = compose_history(data, tI, tR) + f1 = TEST_METRICS['f1'](data, y_pred) + nrmse = TEST_METRICS['nrmse'](data, y_pred) + + metric_rows.append(dict( + dataset = args.dataset, + t_steps = t_steps, + rep = rep, + seed = mcmc_seed, + estimated_pI = bpar.pI, + estimated_pR = bpar.pR, + F1 = f1, + NRMSE = nrmse, + )) + + for row in diag_rows: + trace_rows.append(dict( + dataset = args.dataset, + t_steps = t_steps, + rep = rep, + seed = mcmc_seed, + **row, + )) + + gc.collect() + if getattr(args.device, 'type', None) == 'cuda': + torch.cuda.empty_cache() + + traces = pd.DataFrame(trace_rows, columns = [ + 'dataset', 't_steps', 'rep', 'seed', 'step', 'accept_rate', + 'mean_tI', 'mean_tR', 'mean_tI_avg', 'mean_tR_avg', 'mean_lp' + ]) + metrics = pd.DataFrame(metric_rows, columns = [ + 'dataset', 't_steps', 'rep', 'seed', 'estimated_pI', 'estimated_pR', 'F1', 'NRMSE' + ]) + + traces.to_csv(f'{args.output_prefix}_traces.csv', index = False) + metrics.to_csv(f'{args.output_prefix}_metrics.csv', index = False) + + save_trace_plot( + traces, + y_col = 'accept_rate', + ylabel = 'acceptance rate', + title = f'MCMC acceptance trajectory ({args.dataset})', + fpath = f'{args.output_prefix}_accept.png', + ) + save_trace_plot( + traces, + y_col = 'mean_tI_avg', + ylabel = 'mean infection hitting time', + title = f'MCMC infection-time trace ({args.dataset})', + fpath = f'{args.output_prefix}_mean_tI.png', + ) + save_trace_plot( + traces, + y_col = 'mean_tR_avg', + ylabel = 'mean recovery hitting time', + title = f'MCMC recovery-time trace ({args.dataset})', + fpath = f'{args.output_prefix}_mean_tR.png', + ) + save_metric_plot( + metrics, + metric = 'F1', + ylabel = 'F1', + title = f'Final F1 vs MCMC steps ({args.dataset})', + fpath = f'{args.output_prefix}_F1.png', + ) + save_metric_plot( + metrics, + metric = 'NRMSE', + ylabel = 'NRMSE', + title = f'Final NRMSE vs MCMC steps ({args.dataset})', + fpath = f'{args.output_prefix}_NRMSE.png', + ) + + print(f'[saved] {args.output_prefix}_traces.csv', flush = True) + print(f'[saved] {args.output_prefix}_metrics.csv', flush = True) + print(f'[saved] {args.output_prefix}_accept.png', flush = True) + print(f'[saved] {args.output_prefix}_mean_tI.png', flush = True) + print(f'[saved] {args.output_prefix}_mean_tR.png', flush = True) + print(f'[saved] {args.output_prefix}_F1.png', flush = True) + print(f'[saved] {args.output_prefix}_NRMSE.png', flush = True) + + +if __name__ == '__main__': + main() \ No newline at end of file diff --git a/experiments/exp_nI0_sensitivity.py b/experiments/exp_nI0_sensitivity.py new file mode 100644 index 0000000..c98ade4 --- /dev/null +++ b/experiments/exp_nI0_sensitivity.py @@ -0,0 +1,107 @@ +import os +import sys +import os.path as osp + +ROOT = osp.dirname(osp.dirname(osp.abspath(__file__))) +if ROOT not in sys.path: + sys.path.insert(0, ROOT) + +from inc.data import * +from inc.test import TEST_METRICS +from hermes import run_hermes + +DEFAULT_DATASETS = ['ba-si', 'ba-sir'] +DEFAULT_RATIOS = [0.5, 0.75, 1.0, 1.25, 1.5] + + +def get_args(): + parser = argparse.ArgumentParser() + parser.add_argument('--datasets', type = str, default = ','.join(DEFAULT_DATASETS), help = 'comma-separated dataset names') + parser.add_argument('--ratios', type = str, default = ','.join(map(str, DEFAULT_RATIOS)), help = 'comma-separated I0 ratios') + parser.add_argument('--seed', type = int, default = 12345, help = 'random seed reused across all runs') + parser.add_argument('--data_dir', type = str, required = True, help = 'dataset folder') + parser.add_argument('--output', type = str, required = True, help = 'output csv file name') + parser.add_argument('--device', type = torch.device, default = torch_device(), help = 'torch device') + + # same HERMES hyperparameters as hermes.py + parser.add_argument('--b_pI0', type = float, help = 'initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type = float, help = 'initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type = int, help = 'optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type = float, help = 'learning rate in diffusion parameter estimation') + parser.add_argument('--q_steps', type = int, help = 'training steps for the proposal model') + parser.add_argument('--q_lr', type = float, help = 'learning rate for the proposal model') + parser.add_argument('--q_hid', type = int, help = 'hidden size of the proposal model') + parser.add_argument('--q_gnn', type = int, help = 'number of layers of the GNN in the proposal model') + parser.add_argument('--q_mlp', type = int, help = 'number of layers of the MLP in the proposal model') + parser.add_argument('--q_samples', type = int, help = 'sample size to estimate the loss function of the proposal model') + parser.add_argument('--q_zlim', type = int, help = 'a hyperparameter to stablize gradient') + parser.add_argument('--p_coef', type = float, help = 'the coefficient gamma in the initial distribution P[y_0]') + parser.add_argument('--t_samples', type = int, help = 'MCMC sample size') + parser.add_argument('--t_steps', type = int, help = 'MCMC steps') + parser.add_argument('--t_keep', type = float, help = 'moving average in MCMC') + parser.add_argument('--obs_time', type = str, default = '', help = 'extra observed snapshot times, comma-separated, e.g., 5,7,9') + return parser.parse_args() + + +def parse_list(raw, cast_fn): + return [cast_fn(x.strip()) for x in raw.split(',') if x.strip()] + + +def ratio_to_I0(true_I0, ratio, n_nodes): + assumed_I0 = int(np.floor(true_I0 * ratio + 0.5)) + return max(0, min(int(n_nodes), assumed_I0)) + + +def run_one(data, args, ratio): + true_I0 = resolve_assumed_I0(data, None) + assumed_I0 = ratio_to_I0(true_I0, ratio, data.num_nodes) + if args.seed is not None: + seed_all(args.seed) + + y_pred, extra = run_hermes(data, args, assumed_I0 = assumed_I0, return_extra = True) + + return Dict( + dataset = data.name, + assumed_I0 = assumed_I0, + ratio = ratio, + estimated_pI = extra.pI, + estimated_pR = extra.pR, + F1 = TEST_METRICS['f1'](data, y_pred), + NRMSE = TEST_METRICS['nrmse'](data, y_pred), + ) + + +def main(): + args = get_args() + datasets = parse_list(args.datasets, str) + ratios = parse_list(args.ratios, float) + + rows = [] + for dataset in datasets: + data = data_load(dataset, args.data_dir, args.device) + for ratio in ratios: + print(f'[dataset={dataset}] [ratio={ratio}]', flush = True) + row = run_one(data, args, ratio) + rows.append(dict(row)) + print( + f"[dataset={dataset}] [ratio={ratio}] assumed_I0={row.assumed_I0} " + f"pI={row.estimated_pI:.4f} pR={row.estimated_pR:.4f} " + f"F1={row.F1:.4f} NRMSE={row.NRMSE:.4f}", + flush = True, + ) + gc.collect() + if getattr(args.device, 'type', None) == 'cuda': + torch.cuda.empty_cache() + + cols = ['dataset', 'assumed_I0', 'ratio', 'estimated_pI', 'estimated_pR', 'F1', 'NRMSE'] + df = pd.DataFrame(rows, columns = cols) + + out_dir = osp.dirname(args.output) + if out_dir: + os.makedirs(out_dir, exist_ok = True) + df.to_csv(args.output, index = False) + print(f'[saved] {args.output}', flush = True) + + +if __name__ == '__main__': + main() \ No newline at end of file diff --git a/hermes.py b/hermes.py index fc8e47c..2f02f83 100644 --- a/hermes.py +++ b/hermes.py @@ -27,6 +27,8 @@ def get_args(): parser.add_argument('--t_steps', type = int, help = 'MCMC steps') parser.add_argument('--t_keep', type = float, help = 'moving average in MCMC') parser.add_argument('--obs_time', type = str, default = '', help = 'extra observed snapshot times, comma-separated, e.g., 5,7,9') + parser.add_argument('--assumed_I0', type=int, default=None, + help='assumed initial infected count; default uses the current behavior (true I0 from data)') args = parser.parse_args() return args @@ -274,8 +276,8 @@ def q_loss(q_net, data, I0, bpar, n_samples, obs_time): q_liks, zI0, zR0, zI, zR = q_net.lik_ms(Y=Y, obs_time=obs_time) return -q_liks.mean(), zI0, zR0, zI, zR -def q_train(data, obs_time, bpar, args): - I0 = (data.y[:, 0] == 1).long().sum().item() +def q_train(data, obs_time, bpar, args, assumed_I0 = None): + I0 = resolve_assumed_I0(data, getattr(args, 'assumed_I0', None) if assumed_I0 is None else assumed_I0) q_net = QNet.make(data, obs_time, args) q_net.train() opt = optim.AdamW(q_net.parameters(), lr=args.q_lr) @@ -290,9 +292,12 @@ def q_train(data, obs_time, bpar, args): return q_net @torch.no_grad() -def t_mcmc(data, bpar, q_net, args, obs_time, keepdim=True): +def t_mcmc(data, bpar, q_net, args, obs_time, keepdim=True, assumed_I0=None, diagnostics=False): - I0 = (data.y[:, 0] == 1).long().sum().item() + I0 = resolve_assumed_I0( + data, + getattr(args, 'assumed_I0', None) if assumed_I0 is None else assumed_I0 + ) obs_time = sorted(list(obs_time)) y_obs = torch.stack([data.y[:, t : t + 1] for t in obs_time], dim=2) # (nodes, 1, obs) @@ -304,6 +309,8 @@ def t_mcmc(data, bpar, q_net, args, obs_time, keepdim=True): tI_avg = data_make_t(X, SIR_STATES.I, dim=0).float().mean(dim=1, keepdim=keepdim) tR_avg = data_make_t(X, SIR_STATES.R, dim=0).float().mean(dim=1, keepdim=keepdim) + diag_rows = [] if diagnostics else None + pbar = trange(1, args.t_steps + 1) for step in pbar: Y, lqY = q_net.samp_ms(data.y, zI, zR, args.t_samples, obs_time=obs_time, compute_lik=True) @@ -311,7 +318,9 @@ def t_mcmc(data, bpar, q_net, args, obs_time, keepdim=True): # Hastings acceptance a = torch.rand(args.t_samples, device=args.device) <= torch.exp(lpY + lqX - lpX - lqY) - pbar.set_description(f"[step={step}] acc={a.float().mean().item():.3f}") + acc = a.float().mean().item() + pbar.set_description(f"[step={step}] acc={acc:.3f}") + X = torch.where(a, Y, X) lqX = torch.where(a, lqY, lqX) lpX = torch.where(a, lpY, lpX) @@ -321,31 +330,54 @@ def t_mcmc(data, bpar, q_net, args, obs_time, keepdim=True): tI_avg = args.t_keep * tI_avg + (1.0 - args.t_keep) * tI tR_avg = args.t_keep * tR_avg + (1.0 - args.t_keep) * tR - return tI_avg, tR_avg + if diagnostics: + diag_rows.append(dict( + step=int(step), + accept_rate=float(acc), + mean_tI=float(tI.mean().item()), + mean_tR=float(tR.mean().item()), + mean_tI_avg=float(tI_avg.mean().item()), + mean_tR_avg=float(tR_avg.mean().item()), + mean_lp=float(lpX.mean().item()), + )) -def main(data): + if diagnostics: + return tI_avg, tR_avg, diag_rows + return tI_avg, tR_avg +def run_hermes(data, args, assumed_I0 = None, return_extra = False): # parse obs times obs_time = [int(t) for t in args.obs_time.split(',') if t] obs_time.append(data.T.item()) obs_time = sorted(set(obs_time)) + assumed_I0 = resolve_assumed_I0(data, getattr(args, 'assumed_I0', None) if assumed_I0 is None else assumed_I0) + # estimate diffusion parameters - bpar = b_estim(data, args) + bpar = b_estim(data, args, obs_time = obs_time, assumed_I0 = assumed_I0) print(f'[est] pI={bpar.pI:.4f}, pR={bpar.pR:.4f}', flush = True) + # train a proposal network - q_net = q_train(data, obs_time, bpar, args) + q_net = q_train(data, obs_time, bpar, args, assumed_I0 = assumed_I0) + # estimate transition times - tI, tR = t_mcmc(data, bpar, q_net, args, obs_time = obs_time, keepdim = True) # (nodes, 1) - T = data.T.item() + tI, tR = t_mcmc(data, bpar, q_net, args, obs_time = obs_time, keepdim = True, assumed_I0 = assumed_I0) # (nodes, 1) + tI = tI.round().long() tR = tR.round().long() + # compose a history with torch.no_grad(): y_pred = torch.zeros_like(data.y) # (nodes, T+1) y_pred.scatter_(dim = 1, index = torch.minimum(tI, data.T), src = torch.full_like(tI, 1)) y_pred.scatter_(dim = 1, index = torch.minimum(tR, data.T), src = torch.full_like(tR, 2)) y_pred = y_pred[:, : data.T.item()].cummax(dim = 1).values + if return_extra: + return y_pred, Dict(assumed_I0 = assumed_I0, pI = bpar.pI, pR = bpar.pR) return y_pred -args = get_args() -tester = Tester(args.data_dir, args.device, main) -tester.test([args.dataset], seed = args.seed, rep = 1) +def main(data): + return run_hermes(data, args) + +if __name__ == '__main__': + args = get_args() + tester = Tester(args.data_dir, args.device, main) + tester.test([args.dataset], seed = args.seed, rep = 1) \ No newline at end of file diff --git a/inc/diffus.py b/inc/diffus.py index cc7d7f6..a452857 100644 --- a/inc/diffus.py +++ b/inc/diffus.py @@ -24,12 +24,19 @@ def diffus_sim(edge_index, y0, WI, WR = None): # y0: (nodes, samples); WI: (T, e @torch.no_grad() def diffus_gen(T, n_nodes, edge_index, I0, n_samples, pI, pR): n_edges = edge_index.size(dim = 1) - idx = torch.ones(n_samples, n_nodes, device = edge_index.device).multinomial(I0, replacement = False).T # (I0, samples) y0 = torch.full((n_nodes, n_samples), SIR_STATES.S, dtype = torch.long, device = edge_index.device) - y0.scatter_(index = idx, dim = 0, src = torch.full_like(idx, SIR_STATES.I)) + if I0 > 0: + idx = torch.ones(n_samples, n_nodes, device = edge_index.device).multinomial(I0, replacement = False).T # (I0, samples) + y0.scatter_(index = idx, dim = 0, src = torch.full_like(idx, SIR_STATES.I)) WI = (torch.rand(T, n_edges, n_samples, device = y0.device) < pI).long() WR = (torch.rand(T, n_nodes, n_samples, device = y0.device) < pR).long() if pR > 0 else None return diffus_sim(edge_index, y0, WI, WR) # (T+1, nodes, samples) +def resolve_assumed_I0(data, assumed_I0 = None): + if assumed_I0 is None: + I0 = (data.y[:, 0] == SIR_STATES.I).long().sum().item() + else: + I0 = int(assumed_I0) + return max(0, min(int(data.num_nodes), I0)) def diffus_liks(Y, edge_index, I0, coef, pI, pR): # Y: (T+1, nodes, samples) # assuming Y feasible log1pI = torch_log(1. - pI) if isinstance(pI, torch.Tensor) else math_log(1. - pI) @@ -62,9 +69,9 @@ def __repr__(self, digits = 4): def dict(self): return Dict(pI = self.pI.item(), pR = self.pR.item()) -def _mf_init_from_prior(data, n_nodes, device): +def _mf_init_from_prior(data, n_nodes, device, assumed_I0 = None): # prior: only use I0 count (same as current code) - I0 = (data.y[:, 0] == SIR_STATES.I).sum() + I0 = resolve_assumed_I0(data, assumed_I0) lI = torch.full((n_nodes,), I0 / n_nodes, dtype=torch.float, device=device) lS = torch.full((n_nodes,), 1. - I0 / n_nodes, dtype=torch.float, device=device) lR = torch.zeros(n_nodes, dtype=torch.float, device=device) @@ -77,7 +84,7 @@ def _mf_init_from_snapshot(y, device): lR = (y == SIR_STATES.R).float().to(device) return lS, lI, lR -def b_lik(bpar, data, obs_time=None): +def b_lik(bpar, data, obs_time=None, assumed_I0 = None): device = data.y.device n_nodes = data.num_nodes T = data.T.item() @@ -92,7 +99,7 @@ def b_lik(bpar, data, obs_time=None): obs_time.append(T) # -------- segmented mean-field -------- - lS, lI, lR = _mf_init_from_prior(data, n_nodes, device) + lS, lI, lR = _mf_init_from_prior(data, n_nodes, device, assumed_I0 = assumed_I0) t_prev = 0 lik_total = 0.0 @@ -125,7 +132,7 @@ def b_lik(bpar, data, obs_time=None): lik_total = lik_total / len(obs_time) return lik_total -def b_estim(data, args, obs_time=None): +def b_estim(data, args, obs_time=None, assumed_I0 = None): T = data.T.item() n_cls = data.y[:, T].max().item() + 1 pI, pR = args.b_pI0, (args.b_pR0 if n_cls == 3 else 0.) @@ -136,6 +143,8 @@ def b_estim(data, args, obs_time=None): tmp = [int(t) for t in str(args.obs_time).split(',') if t] tmp.append(T) obs_time = tmp + if assumed_I0 is None and hasattr(args, 'assumed_I0'): + assumed_I0 = args.assumed_I0 bpar = BPar(pI=pI, pR=pR, device=device) bpar.train() @@ -144,7 +153,7 @@ def b_estim(data, args, obs_time=None): for step in pbar: opt.zero_grad() - loss = -b_lik(bpar, data, obs_time=obs_time) + loss = -b_lik(bpar, data, obs_time=obs_time, assumed_I0 = assumed_I0) loss.backward() opt.step() bpar.clamp_() From 76658c58a130203a2fe552e32bdee50c444d351a Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Mon, 6 Apr 2026 21:18:21 -0500 Subject: [PATCH 16/19] Add segmented two-snapshot DITTO baseline --- ditto.py | 29 +++++-- experiments/ditto_seg_twosnap.py | 142 +++++++++++++++++++++++++++++++ 2 files changed, 163 insertions(+), 8 deletions(-) create mode 100644 experiments/ditto_seg_twosnap.py diff --git a/ditto.py b/ditto.py index fb86808..438d8fe 100644 --- a/ditto.py +++ b/ditto.py @@ -226,26 +226,39 @@ def t_mcmc(data, bpar, q_net, args, keepdim = True): tR_avg = args.t_keep * tR_avg + (1. - args.t_keep) * tR # (nodes, 1) return tI_avg, tR_avg # (nodes, 1) -def main(data): +def run_ditto_on_data(data, args): # estimate diffusion parameters bpar = b_estim(data, args) print(f'[est] pI={bpar.pI:.4f}, pR={bpar.pR:.4f}', flush = True) + # train a proposal network q_net = q_train(data, bpar, args) + # estimate transition times - tI, tR = t_mcmc(data, bpar, q_net, args, keepdim = True) # (nodes, 1) - T = data.T.item() + tI, tR = t_mcmc(data, bpar, q_net, args, keepdim = True) # (nodes, 1) tI = tI.round().long() tR = tR.round().long() + # compose a history with torch.no_grad(): - y_pred = torch.zeros_like(data.y) # (nodes, T+1) + y_pred = torch.zeros_like(data.y) # (nodes, T+1) y_pred.scatter_(dim = 1, index = torch.minimum(tI, data.T), src = torch.full_like(tI, 1)) y_pred.scatter_(dim = 1, index = torch.minimum(tR, data.T), src = torch.full_like(tR, 2)) y_pred = y_pred[:, : data.T.item()].cummax(dim = 1).values return y_pred -args = get_args() -tester = Tester(args.data_dir, args.device, main) -tester.test([args.dataset], seed = args.seed, rep = 1) -tester.save(args.output) + +def main(data): + return run_ditto_on_data(data, args) + + +if __name__ == '__main__': + args = get_args() + if args.device is None: + args.device = torch_device() + + tester = Tester(args.data_dir, args.device, main) + tester.test([args.dataset], seed = args.seed, rep = 1) + + if args.output is not None: + tester.save(args.output) \ No newline at end of file diff --git a/experiments/ditto_seg_twosnap.py b/experiments/ditto_seg_twosnap.py new file mode 100644 index 0000000..20b78dc --- /dev/null +++ b/experiments/ditto_seg_twosnap.py @@ -0,0 +1,142 @@ +from inc.diffus import * +from inc.test import * +from ditto import run_ditto_on_data + + +def get_args(): + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type = str, default = None, help = 'single dataset name') + parser.add_argument('--datasets', type = str, default = None, help = 'comma-separated dataset names') + parser.add_argument('--seed', type = int, help = 'random seed') + parser.add_argument('--data_dir', type = str, help = 'dataset folder') + parser.add_argument('--output', type = str, help = 'output file name') + parser.add_argument('--device', type = torch.device, help = 'torch device') + + # segment-and-stitch specific + parser.add_argument('--split_time', type = int, default = None, + help = 'global split time; default=floor(T/2)') + + # same DITTO args as ditto.py + parser.add_argument('--b_pI0', type = float, help = 'initial infection rate in diffusion parameter estimation') + parser.add_argument('--b_pR0', type = float, help = 'initial recovery rate in diffusion parameter estimation') + parser.add_argument('--b_steps', type = int, help = 'optimization steps in diffusion parameter estimation') + parser.add_argument('--b_lr', type = float, help = 'learning rate in diffusion parameter estimation') + + parser.add_argument('--q_steps', type = int, help = 'training steps for the proposal model') + parser.add_argument('--q_lr', type = float, help = 'learning rate for the proposal model') + parser.add_argument('--q_hid', type = int, help = 'hidden size of the proposal model') + parser.add_argument('--q_gnn', type = int, help = 'number of layers of the GNN in the proposal model') + parser.add_argument('--q_mlp', type = int, help = 'number of layers of the MLP in the proposal model') + parser.add_argument('--q_samples', type = int, help = 'sample size to estimate the loss function of the proposal model') + parser.add_argument('--q_zlim', type = int, help = 'a hyperparameter to stabilize gradient') + + parser.add_argument('--p_coef', type = float, help = 'the coefficient gamma in the initial distribution P[y_0]') + parser.add_argument('--t_samples', type = int, help = 'MCMC sample size') + parser.add_argument('--t_steps', type = int, help = 'MCMC steps') + parser.add_argument('--t_keep', type = float, help = 'moving average in MCMC') + + args = parser.parse_args() + + if args.device is None: + args.device = torch_device() + + return args + + +def resolve_datasets(args): + if args.datasets is not None: + datasets = [x.strip() for x in args.datasets.split(',') if x.strip()] + assert len(datasets) > 0, '--datasets is empty' + return datasets + + assert args.dataset is not None, 'either --dataset or --datasets is required' + return [args.dataset] + + +@torch.no_grad() +def make_segment_data(data, t_start, t_end): + """ + Build a local segment subproblem from global times [t_start, t_end]. + + Local time 0 <-> global time t_start + Local time seg_T <-> global time t_end + """ + assert 0 <= t_start < t_end <= data.T.item(), 'invalid segment boundary' + + seg_y = data.y[:, t_start : t_end + 1].detach().clone() # (nodes, seg_T+1) + + seg = Dict( + edge_index = data.edge_index, + num_nodes = data.num_nodes, + y = seg_y, + T = torch.tensor(t_end - t_start, dtype = data.T.dtype, device = data.T.device), + name = f'{data.name}[{t_start},{t_end}]', + ) + return seg + + +@torch.no_grad() +def stitch_two_segments(y1_full, y2_full): + """ + y1_full: global [0, ..., t_split] + y2_full: global [t_split, ..., T] + + Direct stitching rule: + - keep segment 1's terminal snapshot at t_split + - append segment 2 from t_split+1 onward + """ + y_full = torch.cat([y1_full, y2_full[:, 1:]], dim = 1) + return y_full + + +def run_ditto_seg_on_data(data, args): + """ + Naive two-snapshot extension baseline: + 1) run single-snapshot DITTO on [0, t_split], using y_{t_split} as final snapshot + 2) run single-snapshot DITTO on [t_split, T], using y_T as final snapshot + 3) directly stitch the two histories + """ + T = data.T.item() + split_time = args.split_time if args.split_time is not None else (T // 2) + assert 1 <= split_time < T, f'split_time must be in [1, {T - 1}]' + + seg1 = make_segment_data(data, 0, split_time) + seg2 = make_segment_data(data, split_time, T) + + print(f'[split] dataset={data.name} T={T} split={split_time}', flush = True) + + # segment 1: [0, split_time] + print(f'[seg1] run DITTO on [0, {split_time}] with final snapshot y_{split_time}', flush = True) + y1_pred = run_ditto_on_data(seg1, args) # (nodes, split_time) + y1_full = torch.cat([y1_pred, seg1.y[:, -1 :]], dim = 1) # (nodes, split_time+1) + + # segment 2: [split_time, T] + print(f'[seg2] run DITTO on [{split_time}, {T}] with final snapshot y_{T}', flush = True) + y2_pred = run_ditto_on_data(seg2, args) # (nodes, T-split_time) + y2_full = torch.cat([y2_pred, seg2.y[:, -1 :]], dim = 1) # (nodes, T-split_time+1) + + # stitch + y_full = stitch_two_segments(y1_full, y2_full) # (nodes, T+1) + + assert y_full.size(1) == T + 1, 'stitched history has wrong length' + assert torch.equal(y_full[:, split_time], data.y[:, split_time]), 'split snapshot mismatch after stitching' + assert torch.equal(y_full[:, -1], data.y[:, -1]), 'final snapshot mismatch after stitching' + + # Tester expects shape (nodes, T), i.e. all times except the final snapshot + return y_full[:, : T] + + +if __name__ == '__main__': + args = get_args() + datasets = resolve_datasets(args) + + tester = Tester( + args.data_dir, + args.device, + lambda data: run_ditto_seg_on_data(data, args), + ) + + tester.test(datasets, seed = args.seed, rep = 1) + + if args.output is not None: + tester.save(args.output) \ No newline at end of file From 6c4661a04b530b9a6e9df96485c81ddcf15a6491 Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Fri, 10 Apr 2026 00:23:26 -0500 Subject: [PATCH 17/19] Add BA runtime breakdown experiment script --- experiments/runtime_breakdown_ba.py | 484 ++++++++++++++++++++++++++++ 1 file changed, 484 insertions(+) create mode 100644 experiments/runtime_breakdown_ba.py diff --git a/experiments/runtime_breakdown_ba.py b/experiments/runtime_breakdown_ba.py new file mode 100644 index 0000000..f4ee0ca --- /dev/null +++ b/experiments/runtime_breakdown_ba.py @@ -0,0 +1,484 @@ +# experiments/runtime_breakdown_ba.py +# -*- coding: utf-8 -*- + +import os +import sys +import gc +import time +import argparse +import traceback +from collections import OrderedDict + +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import networkx as nx +import pandas as pd +import torch + +# Make the project root importable when this script is placed under experiments/ +ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "")) +if ROOT not in sys.path: + sys.path.insert(0, ROOT) + +from inc.data import data_simulate +from inc.diffus import b_estim +from inc.utils import seed_all +from hermes import q_train, t_mcmc + + +# ---------------------------- +# helpers +# ---------------------------- +def parse_int_list(text): + if text is None: + return [] + text = str(text).strip() + if text == "": + return [] + return [int(x.strip()) for x in text.split(",") if x.strip()] + + +def resolve_device(device_str): + device_str = (device_str or "auto").strip().lower() + if device_str == "auto": + return torch.device("cuda" if torch.cuda.is_available() else "cpu") + if device_str.startswith("cuda") and not torch.cuda.is_available(): + print("[warn] CUDA requested but not available; falling back to CPU.", flush=True) + return torch.device("cpu") + return torch.device(device_str) + + +def cleanup(): + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + +def sync_device(device): + if device.type == "cuda": + torch.cuda.synchronize(device) + + +def timed_call(fn, device): + sync_device(device) + t0 = time.perf_counter() + out = fn() + sync_device(device) + return out, time.perf_counter() - t0 + + +def format_size_label(n): + n = int(n) + if n % 1000 == 0: + return f"{n // 1000}k" + return f"{n:,}" + + +def parse_obs_time(obs_time_str, T): + obs = parse_int_list(obs_time_str) + obs = [int(t) for t in obs if 0 <= int(t) <= int(T)] + if int(T) not in obs: + obs.append(int(T)) + obs = sorted(set(obs)) + return obs + + +# ---------------------------- +# data + compose +# ---------------------------- +def make_ba_sir_data(num_nodes, args): + """ + Create BA-SIR locally in this experiment script, without touching inc/data.py. + Uses the same BA-SIR setting as the current 1k synthetic BA experiment, + except graph size is user-controlled. + """ + Gnx = nx.barabasi_albert_graph( + n=int(num_nodes), + m=int(args.ba_m), + seed=int(args.seed), + ) + + # BA with m>=1 is connected in practice, but keep this for safety. + if not nx.is_connected(Gnx): + Gnx = Gnx.subgraph(max(nx.connected_components(Gnx), key=len)).copy() + + meta = { + "graph_size_actual": int(Gnx.number_of_nodes()), + "num_edges": int(Gnx.number_of_edges()), + } + + params = dict( + fraction_infected=float(args.fraction_infected), + beta=float(args.sim_pI), + gamma=float(args.sim_pR), + ) + + data = data_simulate( + Gnx=Gnx, + seed=int(args.seed), + T=int(args.T), + diffus="sir", + params=params, + ).to(args.device) + + data.name = f"ba-sir-n{meta['graph_size_actual']}" + return data, meta + + +@torch.no_grad() +def compose_history(data, tI, tR): + """ + Same compose logic as run_hermes(), timed separately. + Output shape follows the current project convention: (nodes, T) + and the final observed snapshot is not duplicated here. + """ + tI = tI.round().long() + tR = tR.round().long() + + y_pred = torch.zeros_like(data.y) # (nodes, T+1) + y_pred.scatter_( + dim=1, + index=torch.minimum(tI, data.T), + src=torch.full_like(tI, 1), + ) + y_pred.scatter_( + dim=1, + index=torch.minimum(tR, data.T), + src=torch.full_like(tR, 2), + ) + y_pred = y_pred[:, : data.T.item()].cummax(dim=1).values + return y_pred + + +# ---------------------------- +# one run +# ---------------------------- +def init_row(requested_graph_size, run_order, obs_time): + row = OrderedDict() + row["run_order"] = int(run_order) + row["graph_size_requested"] = int(requested_graph_size) + row["graph_size_actual"] = None + row["num_edges"] = None + row["T"] = None + row["obs_time"] = ",".join(str(t) for t in obs_time) + row["status"] = "pending" + row["fail_stage"] = "" + row["error"] = "" + + row["graph_build_sec"] = None + row["beta_est_sec"] = None + row["proposal_train_sec"] = None + row["mcmc_sec"] = None + row["compose_sec"] = None + row["total_algo_sec"] = None + row["total_wall_sec"] = None + + row["est_pI"] = None + row["est_pR"] = None + return row + + +def run_one_size(requested_graph_size, run_order, args): + obs_time = parse_obs_time(args.obs_time, args.T) + row = init_row( + requested_graph_size=requested_graph_size, + run_order=run_order, + obs_time=obs_time, + ) + + data = None + meta = None + bpar = None + q_net = None + tI = None + tR = None + y_pred = None + stage = "start" + + try: + seed_all(int(args.seed)) + cleanup() + + stage = "graph_build" + (data, meta), graph_build_sec = timed_call( + lambda: make_ba_sir_data(requested_graph_size, args), + args.device, + ) + row["graph_size_actual"] = int(meta["graph_size_actual"]) + row["num_edges"] = int(meta["num_edges"]) + row["T"] = int(data.T.item()) + row["graph_build_sec"] = float(graph_build_sec) + + stage = "beta_est" + bpar, beta_est_sec = timed_call( + lambda: b_estim( + data=data, + args=args, + obs_time=obs_time, + assumed_I0=args.assumed_I0, + ), + args.device, + ) + row["beta_est_sec"] = float(beta_est_sec) + row["est_pI"] = float(bpar.pI) + row["est_pR"] = float(bpar.pR) + + stage = "proposal_train" + q_net, proposal_train_sec = timed_call( + lambda: q_train( + data=data, + obs_time=obs_time, + bpar=bpar, + args=args, + assumed_I0=args.assumed_I0, + ), + args.device, + ) + row["proposal_train_sec"] = float(proposal_train_sec) + + stage = "mcmc" + (tI, tR), mcmc_sec = timed_call( + lambda: t_mcmc( + data=data, + bpar=bpar, + q_net=q_net, + args=args, + obs_time=obs_time, + keepdim=True, + assumed_I0=args.assumed_I0, + diagnostics=False, + ), + args.device, + ) + row["mcmc_sec"] = float(mcmc_sec) + + stage = "compose" + y_pred, compose_sec = timed_call( + lambda: compose_history(data, tI, tR), + args.device, + ) + row["compose_sec"] = float(compose_sec) + + row["total_algo_sec"] = float( + row["beta_est_sec"] + + row["proposal_train_sec"] + + row["mcmc_sec"] + + row["compose_sec"] + ) + row["total_wall_sec"] = float( + row["graph_build_sec"] + row["total_algo_sec"] + ) + row["status"] = "ok" + + except Exception as exc: + row["status"] = "failed" + row["fail_stage"] = stage + row["error"] = f"{type(exc).__name__}: {exc}" + print(f"[failed] n={requested_graph_size}, stage={stage}, error={row['error']}", flush=True) + traceback.print_exc() + + finally: + del data, meta, bpar, q_net, tI, tR, y_pred + cleanup() + + return row + + +# ---------------------------- +# save / plot +# ---------------------------- +def save_rows(rows, csv_path): + df = pd.DataFrame(rows) + df.to_csv(csv_path, index=False) + + +def make_plot(rows, fig_path, q_steps): + df = pd.DataFrame(rows) + df = df[df["status"] == "ok"].copy() + if df.empty: + print(f"[warn] no successful runs; skip figure: {fig_path}", flush=True) + return + + df = df.sort_values("run_order") + + stage_cols = [ + ("beta_est_sec", "beta estimation"), + ("proposal_train_sec", "proposal training"), + ("mcmc_sec", "MCMC"), + ("compose_sec", "compose"), + ] + + x = list(range(len(df))) + xticklabels = [format_size_label(n) for n in df["graph_size_actual"].tolist()] + bottoms = [0.0] * len(df) + + fig, ax = plt.subplots(figsize=(8, 5)) + + for col, label in stage_cols: + vals = df[col].fillna(0.0).astype(float).tolist() + ax.bar(x, vals, bottom=bottoms, label=label) + bottoms = [b + v for b, v in zip(bottoms, vals)] + + totals = df["total_algo_sec"].fillna(0.0).astype(float).tolist() + for i, total in enumerate(totals): + ax.text( + i, + total, + f"{total:.1f}s", + ha="center", + va="bottom", + fontsize=9, + ) + + ax.set_xticks(x) + ax.set_xticklabels(xticklabels) + ax.set_xlabel("BA graph size n") + ax.set_ylabel("running time (sec)") + ax.set_title(f"HERMES runtime breakdown on BA-SIR (q_steps={q_steps})") + ax.legend(frameon=False) + plt.tight_layout() + plt.savefig(fig_path, dpi=200, bbox_inches="tight") + plt.close(fig) + + +# ---------------------------- +# CLI +# ---------------------------- +def get_args(): + parser = argparse.ArgumentParser( + description="Runtime breakdown on BA-SIR for 50k -> 30k -> 1k in one run.", + formatter_class=argparse.ArgumentDefaultsHelpFormatter, + ) + + # run order: do NOT sort, preserve the exact user input order + parser.add_argument( + "--graph_sizes", + type=str, + default="50000,30000,1000", + help="Comma-separated BA sizes; order is preserved exactly.", + ) + parser.add_argument("--seed", type=int, default=123456789) + parser.add_argument("--device", type=str, default="cuda") + parser.add_argument("--output_dir", type=str, default="output/runtime_breakdown_ba") + + # BA-SIR generation + parser.add_argument("--T", type=int, default=10) + parser.add_argument( + "--obs_time", + type=str, + default="5", + help="Extra observed times excluding T; T is appended automatically.", + ) + parser.add_argument("--ba_m", type=int, default=4) + parser.add_argument("--fraction_infected", type=float, default=0.05) + parser.add_argument("--sim_pI", type=float, default=0.1, help="Ground-truth infection rate for simulation.") + parser.add_argument("--sim_pR", type=float, default=0.1, help="Ground-truth recovery rate for simulation.") + + # diffusion parameter estimation + parser.add_argument("--b_pI0", type=float, default=0.001) + parser.add_argument("--b_pR0", type=float, default=0.001) + parser.add_argument("--b_steps", type=int, default=500) + parser.add_argument("--b_lr", type=float, default=0.003) + + # proposal training + parser.add_argument("--q_steps", type=int, default=200) # user-requested change + parser.add_argument("--q_lr", type=float, default=0.003) + parser.add_argument("--q_hid", type=int, default=16) + parser.add_argument("--q_gnn", type=int, default=3) + parser.add_argument("--q_mlp", type=int, default=2) + parser.add_argument("--q_samples", type=int, default=10) + parser.add_argument("--q_zlim", type=int, default=16) + + # MCMC + parser.add_argument("--p_coef", type=float, default=1.0) + parser.add_argument("--t_samples", type=int, default=100) + parser.add_argument("--t_steps", type=int, default=100) + parser.add_argument("--t_keep", type=float, default=0.5) + + # optional initial infected prior + parser.add_argument( + "--assumed_I0", + type=int, + default=None, + help="If None, use the current code behavior (read I0 from data.y[:,0]).", + ) + + args = parser.parse_args() + args.device = resolve_device(args.device) + return args + + +# ---------------------------- +# main +# ---------------------------- +def main(): + args = get_args() + graph_sizes = parse_int_list(args.graph_sizes) + if len(graph_sizes) == 0: + raise ValueError("--graph_sizes is empty.") + + os.makedirs(args.output_dir, exist_ok=True) + + csv_path = os.path.join(args.output_dir, "runtime_breakdown_ba.csv") + fig_path = os.path.join(args.output_dir, "runtime_breakdown_ba.png") + + print("=" * 80, flush=True) + print("BA-SIR runtime breakdown experiment", flush=True) + print(f"device : {args.device}", flush=True) + print(f"graph_sizes : {graph_sizes}", flush=True) + print(f"T : {args.T}", flush=True) + print(f"obs_time : {parse_obs_time(args.obs_time, args.T)}", flush=True) + print(f"q_steps : {args.q_steps}", flush=True) + print("=" * 80, flush=True) + + rows = [] + total_runs = len(graph_sizes) + + for run_order, requested_graph_size in enumerate(graph_sizes, start=1): + print( + f"\n=== [{run_order}/{total_runs}] running BA-SIR with n={requested_graph_size} ===", + flush=True, + ) + row = run_one_size( + requested_graph_size=requested_graph_size, + run_order=run_order, + args=args, + ) + rows.append(row) + save_rows(rows, csv_path) + + if row["status"] == "ok": + print( + "[ok] " + f"n={row['graph_size_actual']:,}, " + f"m={row['num_edges']:,}, " + f"build={row['graph_build_sec']:.2f}s, " + f"beta={row['beta_est_sec']:.2f}s, " + f"q_train={row['proposal_train_sec']:.2f}s, " + f"mcmc={row['mcmc_sec']:.2f}s, " + f"compose={row['compose_sec']:.2f}s, " + f"total_algo={row['total_algo_sec']:.2f}s, " + f"total_wall={row['total_wall_sec']:.2f}s, " + f"est_pI={row['est_pI']:.4f}, " + f"est_pR={row['est_pR']:.4f}", + flush=True, + ) + else: + print( + "[failed] " + f"n={requested_graph_size:,}, " + f"stage={row['fail_stage']}, " + f"error={row['error']}", + flush=True, + ) + + make_plot(rows, fig_path, q_steps=args.q_steps) + + print("\nSaved:") + print(f" CSV : {csv_path}", flush=True) + print(f" FIG : {fig_path}", flush=True) + + +if __name__ == "__main__": + main() \ No newline at end of file From d48900679cfb4e7886336dede57b680ec5614668 Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Fri, 10 Apr 2026 02:09:09 -0500 Subject: [PATCH 18/19] Add BA runtime breakdown experiment script SparseTrain --- .../runtime_breakdown_ba_sparse_rhs.py | 613 ++++++++++++++++++ 1 file changed, 613 insertions(+) create mode 100644 experiments/runtime_breakdown_ba_sparse_rhs.py diff --git a/experiments/runtime_breakdown_ba_sparse_rhs.py b/experiments/runtime_breakdown_ba_sparse_rhs.py new file mode 100644 index 0000000..ac9458d --- /dev/null +++ b/experiments/runtime_breakdown_ba_sparse_rhs.py @@ -0,0 +1,613 @@ +# experiments/runtime_breakdown_ba_sparse_rhs.py +# -*- coding: utf-8 -*- + +import os +import sys +import gc +import time +import argparse +import traceback +from collections import OrderedDict + +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import networkx as nx +import pandas as pd +import torch + +ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) +if ROOT not in sys.path: + sys.path.insert(0, ROOT) + +from inc.data import data_simulate +from inc.diffus import b_estim, SIR_STATES +from inc.utils import seed_all +import hermes as hm + + +# ============================================================ +# monkey patch: only patch q_train hot path (lik_ms / _lik_step) +# ============================================================ +class SparseTrainQNet(hm.QNet): + def _empty_sp(self, size): + idx = torch.empty((2, 0), dtype=torch.long, device=self.device) + val = torch.empty((0,), dtype=self.adj.dtype, device=self.device) + return torch.sparse_coo_tensor( + idx, val, size=size, dtype=self.adj.dtype, device=self.device + ).coalesce() + + def _bool_to_sp(self, mask): + mask = mask.bool() + idx = mask.nonzero(as_tuple=False).T + if idx.numel() == 0: + return self._empty_sp(tuple(mask.shape)) + val = torch.ones(idx.size(1), dtype=self.adj.dtype, device=mask.device) + return torch.sparse_coo_tensor( + idx, + val, + size=tuple(mask.shape), + dtype=self.adj.dtype, + device=mask.device, + ).coalesce() + + def _safe_spmm(self, rhs_sp): + """ + sparse(self.adj) @ sparse(rhs_sp) + If sparse@sparse is not happy on the current backend, fallback safely. + """ + try: + out = torch.sparse.mm(self.adj, rhs_sp) + if out.is_sparse: + return out.coalesce() + return out.to_sparse_coo().coalesce() + except RuntimeError: + out = torch.sparse.mm(self.adj, rhs_sp.to_dense()) + if out.is_sparse: + return out.coalesce() + return out.to_sparse_coo().coalesce() + + def _spmm_bool_dense(self, mask): + rhs_sp = self._bool_to_sp(mask) + out_sp = self._safe_spmm(rhs_sp) + return out_sp.to_dense() + + def _selector_sp(self, rows, active_cols, n_samples): + cols = active_cols.nonzero(as_tuple=False).flatten() + if cols.numel() == 0: + return self._empty_sp((self.n_nodes, n_samples)) + + if torch.is_tensor(rows): + rows = rows.to(device=self.device, dtype=torch.long) + if rows.dim() == 0: + rows = rows.expand(cols.numel()) + else: + rows = rows[cols] + else: + rows = torch.full( + (cols.numel(),), + int(rows), + dtype=torch.long, + device=self.device, + ) + + idx = torch.stack([rows.long(), cols.long()], dim=0) + val = torch.ones(cols.numel(), dtype=self.adj.dtype, device=self.device) + return torch.sparse_coo_tensor( + idx, + val, + size=(self.n_nodes, n_samples), + dtype=self.adj.dtype, + device=self.device, + ).coalesce() + + def _neighbor_delta_dense(self, rows, active_cols, n_samples): + rhs_sp = self._selector_sp(rows, active_cols, n_samples) + if rhs_sp._nnz() == 0: + return torch.zeros( + (self.n_nodes, n_samples), + dtype=torch.long, + device=self.device, + ) + out_sp = self._safe_spmm(rhs_sp) + return (out_sp.to_dense() > 0).to(torch.long) + + def _lik_step(self, y0, y1, lI1, lI0, lR1, lR0, yL=None, reach=None): + # y*, l*, reach: (nodes, samples) + n_nodes, n_samples = y0.shape + lik = self.zero + + # R -> I + msk = (y1 == SIR_STATES.R) + if yL is not None: + msk = msk & (yL != SIR_STATES.R) & reach + lik = lik + torch.where(msk, torch.where(y0 != SIR_STATES.R, lR1, lR0), self.zero) + + # I -> S + uid = lI1.argsort(dim=0, descending=True) + msk = (y1 == SIR_STATES.I) | (msk & (y0 != SIR_STATES.R)) + + rem = torch.where( + msk, + (reach.long() if reach is not None else 1) + self._spmm_bool_dense(msk).long(), + self.n_inf, + ) + + cols = torch.arange(n_samples, dtype=torch.long, device=rem.device) + + for i, u in enumerate(uid): + rem_v = torch.full((self.n_nodes, n_samples), self.n_inf, dtype=rem.dtype, device=rem.device) + rem_v = rem_v.scatter_reduce( + dim=0, + index=self.eidx[0, :, None].expand(-1, n_samples), + src=rem[self.eidx[1]], + reduce="amin", + include_self=True, + ) + rem_v = rem_v.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) + rem_u = rem.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) + + opt = (rem_u > 1) & (rem_v > 1) + if yL is not None: + opt = opt & (yL.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) != SIR_STATES.I) + + msk_u = msk.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) + msk_opt = msk_u & opt + + # keep original semantics + lik = lik + torch.where(msk_opt, torch.where(y0 == SIR_STATES.S, lI1, lI0), self.zero) + + trs = (y0.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) == SIR_STATES.S) + delta = self._neighbor_delta_dense(u, (msk_u & trs), n_samples).to(rem.dtype) + rem = rem - delta + + rem = rem.index_put( + (u, cols), + torch.where( + msk_u, + torch.where(trs, rem_u - 1, self.n_inf), + rem.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0), + ), + ) + msk = msk.index_put((u, cols), msk_opt) + + return lik + + def lik_ms(self, Y, obs_time): + # Y: (T+1, nodes, samples) + assert Y.size(dim=0) == self.T + 1, "lik_ms expects Y with shape (T+1, nodes, samples)" + n_samples = Y.size(dim=2) + K = len(obs_time) + + y_cond = torch.stack([Y[t] for t in obs_time], dim=2) # (nodes, samples, obs) + zI0, zR0, zI, zR = self.forward(y_cond, orig=True) + + zI = zI.clone().detach().requires_grad_(True) + zR = zR.clone().detach().requires_grad_(True) + zI.retain_grad() + zR.retain_grad() + + lik = torch.zeros(n_samples, dtype=zI.dtype, device=self.device) + lI1, lI0 = torch.nn.functional.logsigmoid(zI), torch.nn.functional.logsigmoid(-zI) + lR1, lR0 = torch.nn.functional.logsigmoid(zR), torch.nn.functional.logsigmoid(-zR) + + for i in range(K): + TL, TR = (obs_time[i - 1] if i > 0 else None), obs_time[i] + if TL is None: + for t in range(TR - 1, -1, -1): + lik = lik + self._lik_step(Y[t], Y[t + 1], lI1[t], lI0[t], lR1[t], lR0[t]) + else: + yL = Y[TL] + L = TR - TL + reach = torch.zeros(L, self.n_nodes, n_samples, dtype=torch.bool, device=self.device) + reach[0] = (yL == SIR_STATES.I) + yL_not_R = (yL != SIR_STATES.R) + for d in range(1, L): + nbr = self._spmm_bool_dense(reach[d - 1]) > 0 + reach[d] = reach[d - 1] | (nbr & yL_not_R) + for t in range(TR - 1, TL, -1): + lik = lik + self._lik_step( + Y[t], Y[t + 1], + lI1[t], lI0[t], lR1[t], lR0[t], + yL=yL, + reach=reach[t - TL], + ) + return lik, zI0, zR0, zI, zR + + +# monkey patch +hm.QNet = SparseTrainQNet + + +# ============================================================ +# runtime breakdown experiment +# ============================================================ +def parse_int_list(text): + text = str(text).strip() + if text == "": + return [] + return [int(x.strip()) for x in text.split(",") if x.strip()] + + +def resolve_device(device_str): + device_str = (device_str or "auto").strip().lower() + if device_str == "auto": + return torch.device("cuda" if torch.cuda.is_available() else "cpu") + if device_str.startswith("cuda") and not torch.cuda.is_available(): + print("[warn] CUDA requested but not available; falling back to CPU.", flush=True) + return torch.device("cpu") + return torch.device(device_str) + + +def cleanup(): + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + +def sync_device(device): + if device.type == "cuda": + torch.cuda.synchronize(device) + + +def timed_call(fn, device): + sync_device(device) + t0 = time.perf_counter() + out = fn() + sync_device(device) + return out, time.perf_counter() - t0 + + +def format_size_label(n): + n = int(n) + if n % 1000 == 0: + return f"{n // 1000}k" + return f"{n:,}" + + +def parse_obs_time(obs_time_str, T): + obs = parse_int_list(obs_time_str) + obs = [int(t) for t in obs if 0 <= int(t) <= int(T)] + if int(T) not in obs: + obs.append(int(T)) + return sorted(set(obs)) + + +def make_ba_sir_data(num_nodes, args): + Gnx = nx.barabasi_albert_graph( + n=int(num_nodes), + m=int(args.ba_m), + seed=int(args.seed), + ) + if not nx.is_connected(Gnx): + Gnx = Gnx.subgraph(max(nx.connected_components(Gnx), key=len)).copy() + + meta = { + "graph_size_actual": int(Gnx.number_of_nodes()), + "num_edges": int(Gnx.number_of_edges()), + } + + params = dict( + fraction_infected=float(args.fraction_infected), + beta=float(args.sim_pI), + gamma=float(args.sim_pR), + ) + + data = data_simulate( + Gnx=Gnx, + seed=int(args.seed), + T=int(args.T), + diffus="sir", + params=params, + ).to(args.device) + + data.name = f"ba-sir-n{meta['graph_size_actual']}" + return data, meta + + +@torch.no_grad() +def compose_history(data, tI, tR): + tI = tI.round().long() + tR = tR.round().long() + + y_pred = torch.zeros_like(data.y) + y_pred.scatter_( + dim=1, + index=torch.minimum(tI, data.T), + src=torch.full_like(tI, SIR_STATES.I), + ) + y_pred.scatter_( + dim=1, + index=torch.minimum(tR, data.T), + src=torch.full_like(tR, SIR_STATES.R), + ) + y_pred = y_pred[:, : data.T.item()].cummax(dim=1).values + return y_pred + + +def init_row(requested_graph_size, run_order, obs_time): + row = OrderedDict() + row["run_order"] = int(run_order) + row["graph_size_requested"] = int(requested_graph_size) + row["graph_size_actual"] = None + row["num_edges"] = None + row["T"] = None + row["obs_time"] = ",".join(str(t) for t in obs_time) + row["status"] = "pending" + row["fail_stage"] = "" + row["error"] = "" + + row["graph_build_sec"] = None + row["beta_est_sec"] = None + row["proposal_train_sec"] = None + row["mcmc_sec"] = None + row["compose_sec"] = None + row["total_algo_sec"] = None + row["total_wall_sec"] = None + + row["est_pI"] = None + row["est_pR"] = None + return row + + +def run_one_size(requested_graph_size, run_order, args): + obs_time = parse_obs_time(args.obs_time, args.T) + row = init_row(requested_graph_size, run_order, obs_time) + + data = None + meta = None + bpar = None + q_net = None + tI = None + tR = None + y_pred = None + stage = "start" + + try: + seed_all(int(args.seed)) + cleanup() + + stage = "graph_build" + (data, meta), graph_build_sec = timed_call( + lambda: make_ba_sir_data(requested_graph_size, args), + args.device, + ) + row["graph_size_actual"] = int(meta["graph_size_actual"]) + row["num_edges"] = int(meta["num_edges"]) + row["T"] = int(data.T.item()) + row["graph_build_sec"] = float(graph_build_sec) + + stage = "beta_est" + bpar, beta_est_sec = timed_call( + lambda: b_estim( + data=data, + args=args, + obs_time=obs_time, + assumed_I0=args.assumed_I0, + ), + args.device, + ) + row["beta_est_sec"] = float(beta_est_sec) + row["est_pI"] = float(bpar.pI) + row["est_pR"] = float(bpar.pR) + + stage = "proposal_train" + q_net, proposal_train_sec = timed_call( + lambda: hm.q_train( + data=data, + obs_time=obs_time, + bpar=bpar, + args=args, + assumed_I0=args.assumed_I0, + ), + args.device, + ) + row["proposal_train_sec"] = float(proposal_train_sec) + + stage = "mcmc" + (tI, tR), mcmc_sec = timed_call( + lambda: hm.t_mcmc( + data=data, + bpar=bpar, + q_net=q_net, + args=args, + obs_time=obs_time, + keepdim=True, + assumed_I0=args.assumed_I0, + diagnostics=False, + ), + args.device, + ) + row["mcmc_sec"] = float(mcmc_sec) + + stage = "compose" + y_pred, compose_sec = timed_call( + lambda: compose_history(data, tI, tR), + args.device, + ) + row["compose_sec"] = float(compose_sec) + + row["total_algo_sec"] = float( + row["beta_est_sec"] + + row["proposal_train_sec"] + + row["mcmc_sec"] + + row["compose_sec"] + ) + row["total_wall_sec"] = float(row["graph_build_sec"] + row["total_algo_sec"]) + row["status"] = "ok" + + except Exception as exc: + row["status"] = "failed" + row["fail_stage"] = stage + row["error"] = f"{type(exc).__name__}: {exc}" + print(f"[failed] n={requested_graph_size}, stage={stage}, error={row['error']}", flush=True) + traceback.print_exc() + + finally: + del data, meta, bpar, q_net, tI, tR, y_pred + cleanup() + + return row + + +def save_rows(rows, csv_path): + df = pd.DataFrame(rows) + df.to_csv(csv_path, index=False) + + +def make_plot(rows, fig_path, q_steps): + df = pd.DataFrame(rows) + df = df[df["status"] == "ok"].copy() + if df.empty: + print(f"[warn] no successful runs; skip figure: {fig_path}", flush=True) + return + + df = df.sort_values("run_order") + + stage_cols = [ + ("beta_est_sec", "beta estimation"), + ("proposal_train_sec", "proposal training"), + ("mcmc_sec", "MCMC"), + ("compose_sec", "compose"), + ] + + x = list(range(len(df))) + xticklabels = [format_size_label(n) for n in df["graph_size_actual"].tolist()] + bottoms = [0.0] * len(df) + + fig, ax = plt.subplots(figsize=(8, 5)) + + for col, label in stage_cols: + vals = df[col].fillna(0.0).astype(float).tolist() + ax.bar(x, vals, bottom=bottoms, label=label) + bottoms = [b + v for b, v in zip(bottoms, vals)] + + totals = df["total_algo_sec"].fillna(0.0).astype(float).tolist() + for i, total in enumerate(totals): + ax.text(i, total, f"{total:.1f}s", ha="center", va="bottom", fontsize=9) + + ax.set_xticks(x) + ax.set_xticklabels(xticklabels) + ax.set_xlabel("BA graph size n") + ax.set_ylabel("running time (sec)") + ax.set_title(f"HERMES runtime breakdown on BA-SIR (sparse RHS, q_steps={q_steps})") + ax.legend(frameon=False) + plt.tight_layout() + plt.savefig(fig_path, dpi=200, bbox_inches="tight") + plt.close(fig) + + +def get_args(): + parser = argparse.ArgumentParser( + description="Runtime breakdown on BA-SIR with sparse RHS patch for proposal training.", + formatter_class=argparse.ArgumentDefaultsHelpFormatter, + ) + + parser.add_argument("--graph_sizes", type=str, default="50000,30000,1000") + parser.add_argument("--seed", type=int, default=123456789) + parser.add_argument("--device", type=str, default="cuda") + parser.add_argument("--output_dir", type=str, default="output/runtime_breakdown_ba_sparse_rhs") + + # BA-SIR generation + parser.add_argument("--T", type=int, default=10) + parser.add_argument("--obs_time", type=str, default="5") + parser.add_argument("--ba_m", type=int, default=4) + parser.add_argument("--fraction_infected", type=float, default=0.05) + parser.add_argument("--sim_pI", type=float, default=0.1) + parser.add_argument("--sim_pR", type=float, default=0.1) + + # diffusion parameter estimation + parser.add_argument("--b_pI0", type=float, default=0.001) + parser.add_argument("--b_pR0", type=float, default=0.001) + parser.add_argument("--b_steps", type=int, default=500) + parser.add_argument("--b_lr", type=float, default=0.003) + + # proposal training + parser.add_argument("--q_steps", type=int, default=200) + parser.add_argument("--q_lr", type=float, default=0.003) + parser.add_argument("--q_hid", type=int, default=16) + parser.add_argument("--q_gnn", type=int, default=3) + parser.add_argument("--q_mlp", type=int, default=2) + parser.add_argument("--q_samples", type=int, default=10) + parser.add_argument("--q_zlim", type=int, default=16) + + # MCMC + parser.add_argument("--p_coef", type=float, default=1.0) + parser.add_argument("--t_samples", type=int, default=100) + parser.add_argument("--t_steps", type=int, default=100) + parser.add_argument("--t_keep", type=float, default=0.5) + + parser.add_argument("--assumed_I0", type=int, default=None) + + args = parser.parse_args() + args.device = resolve_device(args.device) + return args + + +def main(): + args = get_args() + graph_sizes = parse_int_list(args.graph_sizes) + if len(graph_sizes) == 0: + raise ValueError("--graph_sizes is empty.") + + os.makedirs(args.output_dir, exist_ok=True) + + csv_path = os.path.join(args.output_dir, "runtime_breakdown_ba_sparse_rhs.csv") + fig_path = os.path.join(args.output_dir, "runtime_breakdown_ba_sparse_rhs.png") + + print("=" * 80, flush=True) + print("BA-SIR runtime breakdown experiment (sparse RHS patch)", flush=True) + print(f"device : {args.device}", flush=True) + print(f"graph_sizes : {graph_sizes}", flush=True) + print(f"T : {args.T}", flush=True) + print(f"obs_time : {parse_obs_time(args.obs_time, args.T)}", flush=True) + print(f"q_steps : {args.q_steps}", flush=True) + print("=" * 80, flush=True) + + rows = [] + total_runs = len(graph_sizes) + + for run_order, requested_graph_size in enumerate(graph_sizes, start=1): + print(f"\n=== [{run_order}/{total_runs}] running BA-SIR with n={requested_graph_size} ===", flush=True) + row = run_one_size( + requested_graph_size=requested_graph_size, + run_order=run_order, + args=args, + ) + rows.append(row) + save_rows(rows, csv_path) + + if row["status"] == "ok": + print( + "[ok] " + f"n={row['graph_size_actual']:,}, " + f"m={row['num_edges']:,}, " + f"build={row['graph_build_sec']:.2f}s, " + f"beta={row['beta_est_sec']:.2f}s, " + f"q_train={row['proposal_train_sec']:.2f}s, " + f"mcmc={row['mcmc_sec']:.2f}s, " + f"compose={row['compose_sec']:.2f}s, " + f"total_algo={row['total_algo_sec']:.2f}s, " + f"total_wall={row['total_wall_sec']:.2f}s, " + f"est_pI={row['est_pI']:.4f}, " + f"est_pR={row['est_pR']:.4f}", + flush=True, + ) + else: + print( + "[failed] " + f"n={requested_graph_size:,}, " + f"stage={row['fail_stage']}, " + f"error={row['error']}", + flush=True, + ) + + make_plot(rows, fig_path, q_steps=args.q_steps) + + print("\nSaved:") + print(f" CSV : {csv_path}", flush=True) + print(f" FIG : {fig_path}", flush=True) + + +if __name__ == "__main__": + main() \ No newline at end of file From 5679d14bf5f7dd5922088c44310befe5a7448b9e Mon Sep 17 00:00:00 2001 From: Yijing Zuo Date: Fri, 10 Apr 2026 11:18:04 -0500 Subject: [PATCH 19/19] Add sparse runtime breakdown experiment script --- experiments/runtime_breakdown_ba_sparse.py | 160 +++++ .../runtime_breakdown_ba_sparse_rhs.py | 613 ------------------ 2 files changed, 160 insertions(+), 613 deletions(-) create mode 100644 experiments/runtime_breakdown_ba_sparse.py delete mode 100644 experiments/runtime_breakdown_ba_sparse_rhs.py diff --git a/experiments/runtime_breakdown_ba_sparse.py b/experiments/runtime_breakdown_ba_sparse.py new file mode 100644 index 0000000..cf1286d --- /dev/null +++ b/experiments/runtime_breakdown_ba_sparse.py @@ -0,0 +1,160 @@ +# experiments/runtime_breakdown_ba_sparse.py +# -*- coding: utf-8 -*- + +import os +import sys +import torch + +THIS_DIR = os.path.dirname(os.path.abspath(__file__)) +ROOT = os.path.abspath(os.path.join(THIS_DIR, "..")) + +if THIS_DIR not in sys.path: + sys.path.insert(0, THIS_DIR) +if ROOT not in sys.path: + sys.path.insert(0, ROOT) + +import runtime_breakdown_ba as base +from hermes import QNet +from inc.diffus import SIR_STATES + + +def apply_sparse_rhs_patch(): + @torch.no_grad() + def _lik_step_sparse_rhs(self, y0, y1, lI1, lI0, lR1, lR0, yL=None, reach=None): + n_nodes, n_samples = y0.shape + lik = self.zero + + # R -> I part (unchanged) + msk = (y1 == SIR_STATES.R) + if yL is not None: + msk = msk & (yL != SIR_STATES.R) & reach + lik = lik + torch.where( + msk, + torch.where(y0 != SIR_STATES.R, lR1, lR0), + self.zero, + ) + + # I -> S part + uid = lI1.argsort(dim=0, descending=True) # (nodes, samples) + msk = (y1 == SIR_STATES.I) | (msk & (y0 != SIR_STATES.R)) + + rem = torch.where( + msk, + (reach.long() if reach is not None else 1) + + torch.sparse.mm(self.adj, msk.float()).long(), + self.n_inf, + ) # (nodes, samples) + + cols = torch.arange(n_samples, dtype=torch.long, device=rem.device) + + for i, u in enumerate(uid): + rem_v = torch.full( + (self.n_nodes, n_samples), + self.n_inf, + dtype=rem.dtype, + device=rem.device, + ) + rem_v = rem_v.scatter_reduce( + dim=0, + index=self.eidx[0, :, None].expand(-1, n_samples), + src=rem[self.eidx[1]], + reduce="amin", + include_self=True, + ) + rem_v = rem_v.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) # (samples,) + rem_u = rem.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) # (samples,) + + opt = (rem_u > 1) & (rem_v > 1) + if yL is not None: + opt = opt & ( + yL.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) != SIR_STATES.I + ) + + msk_u = msk.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) # (samples,) + msk_opt = msk_u & opt + + lik = lik + torch.where( + msk_opt, + torch.where(y0 == SIR_STATES.S, lI1, lI0), + self.zero, + ) + + trs = ( + y0.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) + == SIR_STATES.S + ) # (samples,) + + # ========================================================== + # OLD: + # rem = rem - (torch.sparse.mm( + # self.adj, + # torch.zeros(rem.size(), ...).scatter(...) + # ) > 0).to(rem.dtype) + # + # NEW: + # build RHS directly as sparse COO and use sparse indices. + # ========================================================== + active_cols = torch.nonzero(msk_u & trs, as_tuple=False).flatten() + if active_cols.numel() > 0: + rhs_row = u.index_select(0, active_cols) # row index varies by sample + rhs_col = active_cols + rhs_idx = torch.stack([rhs_row, rhs_col], dim=0) + rhs_val = torch.ones( + active_cols.numel(), + dtype=self.adj.dtype, + device=rem.device, + ) + + rhs = torch.sparse_coo_tensor( + indices=rhs_idx, + values=rhs_val, + size=rem.size(), # (nodes, samples) + dtype=self.adj.dtype, + device=rem.device, + ).coalesce() + + nbr_hit = torch.sparse.mm(self.adj, rhs) + if nbr_hit.layout != torch.sparse_coo: + nbr_hit = nbr_hit.to_sparse_coo() + nbr_hit = nbr_hit.coalesce() + + # IMPORTANT: + # sparse tensor usually should not continue with `> 0` here. + # Use sparse indices directly. + if nbr_hit._nnz() > 0: + hit_idx = nbr_hit.indices() + rem.index_put_( + (hit_idx[0], hit_idx[1]), + -torch.ones( + hit_idx.size(1), + dtype=rem.dtype, + device=rem.device, + ), + accumulate=True, + ) + + # self-row update (unchanged logic; in-place for less allocation) + rem.index_put_( + (u, cols), + torch.where( + msk_u, + torch.where(trs, rem_u - 1, self.n_inf), + rem.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0), + ), + accumulate=False, + ) + msk.index_put_((u, cols), msk_opt, accumulate=False) + + return lik + + QNet._lik_step = _lik_step_sparse_rhs + print("[patch] QNet._lik_step -> sparse RHS version", flush=True) + + +def main(): + apply_sparse_rhs_patch() + base.main() + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/experiments/runtime_breakdown_ba_sparse_rhs.py b/experiments/runtime_breakdown_ba_sparse_rhs.py deleted file mode 100644 index ac9458d..0000000 --- a/experiments/runtime_breakdown_ba_sparse_rhs.py +++ /dev/null @@ -1,613 +0,0 @@ -# experiments/runtime_breakdown_ba_sparse_rhs.py -# -*- coding: utf-8 -*- - -import os -import sys -import gc -import time -import argparse -import traceback -from collections import OrderedDict - -import matplotlib -matplotlib.use("Agg") -import matplotlib.pyplot as plt -import networkx as nx -import pandas as pd -import torch - -ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) -if ROOT not in sys.path: - sys.path.insert(0, ROOT) - -from inc.data import data_simulate -from inc.diffus import b_estim, SIR_STATES -from inc.utils import seed_all -import hermes as hm - - -# ============================================================ -# monkey patch: only patch q_train hot path (lik_ms / _lik_step) -# ============================================================ -class SparseTrainQNet(hm.QNet): - def _empty_sp(self, size): - idx = torch.empty((2, 0), dtype=torch.long, device=self.device) - val = torch.empty((0,), dtype=self.adj.dtype, device=self.device) - return torch.sparse_coo_tensor( - idx, val, size=size, dtype=self.adj.dtype, device=self.device - ).coalesce() - - def _bool_to_sp(self, mask): - mask = mask.bool() - idx = mask.nonzero(as_tuple=False).T - if idx.numel() == 0: - return self._empty_sp(tuple(mask.shape)) - val = torch.ones(idx.size(1), dtype=self.adj.dtype, device=mask.device) - return torch.sparse_coo_tensor( - idx, - val, - size=tuple(mask.shape), - dtype=self.adj.dtype, - device=mask.device, - ).coalesce() - - def _safe_spmm(self, rhs_sp): - """ - sparse(self.adj) @ sparse(rhs_sp) - If sparse@sparse is not happy on the current backend, fallback safely. - """ - try: - out = torch.sparse.mm(self.adj, rhs_sp) - if out.is_sparse: - return out.coalesce() - return out.to_sparse_coo().coalesce() - except RuntimeError: - out = torch.sparse.mm(self.adj, rhs_sp.to_dense()) - if out.is_sparse: - return out.coalesce() - return out.to_sparse_coo().coalesce() - - def _spmm_bool_dense(self, mask): - rhs_sp = self._bool_to_sp(mask) - out_sp = self._safe_spmm(rhs_sp) - return out_sp.to_dense() - - def _selector_sp(self, rows, active_cols, n_samples): - cols = active_cols.nonzero(as_tuple=False).flatten() - if cols.numel() == 0: - return self._empty_sp((self.n_nodes, n_samples)) - - if torch.is_tensor(rows): - rows = rows.to(device=self.device, dtype=torch.long) - if rows.dim() == 0: - rows = rows.expand(cols.numel()) - else: - rows = rows[cols] - else: - rows = torch.full( - (cols.numel(),), - int(rows), - dtype=torch.long, - device=self.device, - ) - - idx = torch.stack([rows.long(), cols.long()], dim=0) - val = torch.ones(cols.numel(), dtype=self.adj.dtype, device=self.device) - return torch.sparse_coo_tensor( - idx, - val, - size=(self.n_nodes, n_samples), - dtype=self.adj.dtype, - device=self.device, - ).coalesce() - - def _neighbor_delta_dense(self, rows, active_cols, n_samples): - rhs_sp = self._selector_sp(rows, active_cols, n_samples) - if rhs_sp._nnz() == 0: - return torch.zeros( - (self.n_nodes, n_samples), - dtype=torch.long, - device=self.device, - ) - out_sp = self._safe_spmm(rhs_sp) - return (out_sp.to_dense() > 0).to(torch.long) - - def _lik_step(self, y0, y1, lI1, lI0, lR1, lR0, yL=None, reach=None): - # y*, l*, reach: (nodes, samples) - n_nodes, n_samples = y0.shape - lik = self.zero - - # R -> I - msk = (y1 == SIR_STATES.R) - if yL is not None: - msk = msk & (yL != SIR_STATES.R) & reach - lik = lik + torch.where(msk, torch.where(y0 != SIR_STATES.R, lR1, lR0), self.zero) - - # I -> S - uid = lI1.argsort(dim=0, descending=True) - msk = (y1 == SIR_STATES.I) | (msk & (y0 != SIR_STATES.R)) - - rem = torch.where( - msk, - (reach.long() if reach is not None else 1) + self._spmm_bool_dense(msk).long(), - self.n_inf, - ) - - cols = torch.arange(n_samples, dtype=torch.long, device=rem.device) - - for i, u in enumerate(uid): - rem_v = torch.full((self.n_nodes, n_samples), self.n_inf, dtype=rem.dtype, device=rem.device) - rem_v = rem_v.scatter_reduce( - dim=0, - index=self.eidx[0, :, None].expand(-1, n_samples), - src=rem[self.eidx[1]], - reduce="amin", - include_self=True, - ) - rem_v = rem_v.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) - rem_u = rem.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) - - opt = (rem_u > 1) & (rem_v > 1) - if yL is not None: - opt = opt & (yL.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) != SIR_STATES.I) - - msk_u = msk.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) - msk_opt = msk_u & opt - - # keep original semantics - lik = lik + torch.where(msk_opt, torch.where(y0 == SIR_STATES.S, lI1, lI0), self.zero) - - trs = (y0.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0) == SIR_STATES.S) - delta = self._neighbor_delta_dense(u, (msk_u & trs), n_samples).to(rem.dtype) - rem = rem - delta - - rem = rem.index_put( - (u, cols), - torch.where( - msk_u, - torch.where(trs, rem_u - 1, self.n_inf), - rem.gather(dim=0, index=u.unsqueeze(dim=0)).squeeze(dim=0), - ), - ) - msk = msk.index_put((u, cols), msk_opt) - - return lik - - def lik_ms(self, Y, obs_time): - # Y: (T+1, nodes, samples) - assert Y.size(dim=0) == self.T + 1, "lik_ms expects Y with shape (T+1, nodes, samples)" - n_samples = Y.size(dim=2) - K = len(obs_time) - - y_cond = torch.stack([Y[t] for t in obs_time], dim=2) # (nodes, samples, obs) - zI0, zR0, zI, zR = self.forward(y_cond, orig=True) - - zI = zI.clone().detach().requires_grad_(True) - zR = zR.clone().detach().requires_grad_(True) - zI.retain_grad() - zR.retain_grad() - - lik = torch.zeros(n_samples, dtype=zI.dtype, device=self.device) - lI1, lI0 = torch.nn.functional.logsigmoid(zI), torch.nn.functional.logsigmoid(-zI) - lR1, lR0 = torch.nn.functional.logsigmoid(zR), torch.nn.functional.logsigmoid(-zR) - - for i in range(K): - TL, TR = (obs_time[i - 1] if i > 0 else None), obs_time[i] - if TL is None: - for t in range(TR - 1, -1, -1): - lik = lik + self._lik_step(Y[t], Y[t + 1], lI1[t], lI0[t], lR1[t], lR0[t]) - else: - yL = Y[TL] - L = TR - TL - reach = torch.zeros(L, self.n_nodes, n_samples, dtype=torch.bool, device=self.device) - reach[0] = (yL == SIR_STATES.I) - yL_not_R = (yL != SIR_STATES.R) - for d in range(1, L): - nbr = self._spmm_bool_dense(reach[d - 1]) > 0 - reach[d] = reach[d - 1] | (nbr & yL_not_R) - for t in range(TR - 1, TL, -1): - lik = lik + self._lik_step( - Y[t], Y[t + 1], - lI1[t], lI0[t], lR1[t], lR0[t], - yL=yL, - reach=reach[t - TL], - ) - return lik, zI0, zR0, zI, zR - - -# monkey patch -hm.QNet = SparseTrainQNet - - -# ============================================================ -# runtime breakdown experiment -# ============================================================ -def parse_int_list(text): - text = str(text).strip() - if text == "": - return [] - return [int(x.strip()) for x in text.split(",") if x.strip()] - - -def resolve_device(device_str): - device_str = (device_str or "auto").strip().lower() - if device_str == "auto": - return torch.device("cuda" if torch.cuda.is_available() else "cpu") - if device_str.startswith("cuda") and not torch.cuda.is_available(): - print("[warn] CUDA requested but not available; falling back to CPU.", flush=True) - return torch.device("cpu") - return torch.device(device_str) - - -def cleanup(): - gc.collect() - if torch.cuda.is_available(): - torch.cuda.empty_cache() - - -def sync_device(device): - if device.type == "cuda": - torch.cuda.synchronize(device) - - -def timed_call(fn, device): - sync_device(device) - t0 = time.perf_counter() - out = fn() - sync_device(device) - return out, time.perf_counter() - t0 - - -def format_size_label(n): - n = int(n) - if n % 1000 == 0: - return f"{n // 1000}k" - return f"{n:,}" - - -def parse_obs_time(obs_time_str, T): - obs = parse_int_list(obs_time_str) - obs = [int(t) for t in obs if 0 <= int(t) <= int(T)] - if int(T) not in obs: - obs.append(int(T)) - return sorted(set(obs)) - - -def make_ba_sir_data(num_nodes, args): - Gnx = nx.barabasi_albert_graph( - n=int(num_nodes), - m=int(args.ba_m), - seed=int(args.seed), - ) - if not nx.is_connected(Gnx): - Gnx = Gnx.subgraph(max(nx.connected_components(Gnx), key=len)).copy() - - meta = { - "graph_size_actual": int(Gnx.number_of_nodes()), - "num_edges": int(Gnx.number_of_edges()), - } - - params = dict( - fraction_infected=float(args.fraction_infected), - beta=float(args.sim_pI), - gamma=float(args.sim_pR), - ) - - data = data_simulate( - Gnx=Gnx, - seed=int(args.seed), - T=int(args.T), - diffus="sir", - params=params, - ).to(args.device) - - data.name = f"ba-sir-n{meta['graph_size_actual']}" - return data, meta - - -@torch.no_grad() -def compose_history(data, tI, tR): - tI = tI.round().long() - tR = tR.round().long() - - y_pred = torch.zeros_like(data.y) - y_pred.scatter_( - dim=1, - index=torch.minimum(tI, data.T), - src=torch.full_like(tI, SIR_STATES.I), - ) - y_pred.scatter_( - dim=1, - index=torch.minimum(tR, data.T), - src=torch.full_like(tR, SIR_STATES.R), - ) - y_pred = y_pred[:, : data.T.item()].cummax(dim=1).values - return y_pred - - -def init_row(requested_graph_size, run_order, obs_time): - row = OrderedDict() - row["run_order"] = int(run_order) - row["graph_size_requested"] = int(requested_graph_size) - row["graph_size_actual"] = None - row["num_edges"] = None - row["T"] = None - row["obs_time"] = ",".join(str(t) for t in obs_time) - row["status"] = "pending" - row["fail_stage"] = "" - row["error"] = "" - - row["graph_build_sec"] = None - row["beta_est_sec"] = None - row["proposal_train_sec"] = None - row["mcmc_sec"] = None - row["compose_sec"] = None - row["total_algo_sec"] = None - row["total_wall_sec"] = None - - row["est_pI"] = None - row["est_pR"] = None - return row - - -def run_one_size(requested_graph_size, run_order, args): - obs_time = parse_obs_time(args.obs_time, args.T) - row = init_row(requested_graph_size, run_order, obs_time) - - data = None - meta = None - bpar = None - q_net = None - tI = None - tR = None - y_pred = None - stage = "start" - - try: - seed_all(int(args.seed)) - cleanup() - - stage = "graph_build" - (data, meta), graph_build_sec = timed_call( - lambda: make_ba_sir_data(requested_graph_size, args), - args.device, - ) - row["graph_size_actual"] = int(meta["graph_size_actual"]) - row["num_edges"] = int(meta["num_edges"]) - row["T"] = int(data.T.item()) - row["graph_build_sec"] = float(graph_build_sec) - - stage = "beta_est" - bpar, beta_est_sec = timed_call( - lambda: b_estim( - data=data, - args=args, - obs_time=obs_time, - assumed_I0=args.assumed_I0, - ), - args.device, - ) - row["beta_est_sec"] = float(beta_est_sec) - row["est_pI"] = float(bpar.pI) - row["est_pR"] = float(bpar.pR) - - stage = "proposal_train" - q_net, proposal_train_sec = timed_call( - lambda: hm.q_train( - data=data, - obs_time=obs_time, - bpar=bpar, - args=args, - assumed_I0=args.assumed_I0, - ), - args.device, - ) - row["proposal_train_sec"] = float(proposal_train_sec) - - stage = "mcmc" - (tI, tR), mcmc_sec = timed_call( - lambda: hm.t_mcmc( - data=data, - bpar=bpar, - q_net=q_net, - args=args, - obs_time=obs_time, - keepdim=True, - assumed_I0=args.assumed_I0, - diagnostics=False, - ), - args.device, - ) - row["mcmc_sec"] = float(mcmc_sec) - - stage = "compose" - y_pred, compose_sec = timed_call( - lambda: compose_history(data, tI, tR), - args.device, - ) - row["compose_sec"] = float(compose_sec) - - row["total_algo_sec"] = float( - row["beta_est_sec"] - + row["proposal_train_sec"] - + row["mcmc_sec"] - + row["compose_sec"] - ) - row["total_wall_sec"] = float(row["graph_build_sec"] + row["total_algo_sec"]) - row["status"] = "ok" - - except Exception as exc: - row["status"] = "failed" - row["fail_stage"] = stage - row["error"] = f"{type(exc).__name__}: {exc}" - print(f"[failed] n={requested_graph_size}, stage={stage}, error={row['error']}", flush=True) - traceback.print_exc() - - finally: - del data, meta, bpar, q_net, tI, tR, y_pred - cleanup() - - return row - - -def save_rows(rows, csv_path): - df = pd.DataFrame(rows) - df.to_csv(csv_path, index=False) - - -def make_plot(rows, fig_path, q_steps): - df = pd.DataFrame(rows) - df = df[df["status"] == "ok"].copy() - if df.empty: - print(f"[warn] no successful runs; skip figure: {fig_path}", flush=True) - return - - df = df.sort_values("run_order") - - stage_cols = [ - ("beta_est_sec", "beta estimation"), - ("proposal_train_sec", "proposal training"), - ("mcmc_sec", "MCMC"), - ("compose_sec", "compose"), - ] - - x = list(range(len(df))) - xticklabels = [format_size_label(n) for n in df["graph_size_actual"].tolist()] - bottoms = [0.0] * len(df) - - fig, ax = plt.subplots(figsize=(8, 5)) - - for col, label in stage_cols: - vals = df[col].fillna(0.0).astype(float).tolist() - ax.bar(x, vals, bottom=bottoms, label=label) - bottoms = [b + v for b, v in zip(bottoms, vals)] - - totals = df["total_algo_sec"].fillna(0.0).astype(float).tolist() - for i, total in enumerate(totals): - ax.text(i, total, f"{total:.1f}s", ha="center", va="bottom", fontsize=9) - - ax.set_xticks(x) - ax.set_xticklabels(xticklabels) - ax.set_xlabel("BA graph size n") - ax.set_ylabel("running time (sec)") - ax.set_title(f"HERMES runtime breakdown on BA-SIR (sparse RHS, q_steps={q_steps})") - ax.legend(frameon=False) - plt.tight_layout() - plt.savefig(fig_path, dpi=200, bbox_inches="tight") - plt.close(fig) - - -def get_args(): - parser = argparse.ArgumentParser( - description="Runtime breakdown on BA-SIR with sparse RHS patch for proposal training.", - formatter_class=argparse.ArgumentDefaultsHelpFormatter, - ) - - parser.add_argument("--graph_sizes", type=str, default="50000,30000,1000") - parser.add_argument("--seed", type=int, default=123456789) - parser.add_argument("--device", type=str, default="cuda") - parser.add_argument("--output_dir", type=str, default="output/runtime_breakdown_ba_sparse_rhs") - - # BA-SIR generation - parser.add_argument("--T", type=int, default=10) - parser.add_argument("--obs_time", type=str, default="5") - parser.add_argument("--ba_m", type=int, default=4) - parser.add_argument("--fraction_infected", type=float, default=0.05) - parser.add_argument("--sim_pI", type=float, default=0.1) - parser.add_argument("--sim_pR", type=float, default=0.1) - - # diffusion parameter estimation - parser.add_argument("--b_pI0", type=float, default=0.001) - parser.add_argument("--b_pR0", type=float, default=0.001) - parser.add_argument("--b_steps", type=int, default=500) - parser.add_argument("--b_lr", type=float, default=0.003) - - # proposal training - parser.add_argument("--q_steps", type=int, default=200) - parser.add_argument("--q_lr", type=float, default=0.003) - parser.add_argument("--q_hid", type=int, default=16) - parser.add_argument("--q_gnn", type=int, default=3) - parser.add_argument("--q_mlp", type=int, default=2) - parser.add_argument("--q_samples", type=int, default=10) - parser.add_argument("--q_zlim", type=int, default=16) - - # MCMC - parser.add_argument("--p_coef", type=float, default=1.0) - parser.add_argument("--t_samples", type=int, default=100) - parser.add_argument("--t_steps", type=int, default=100) - parser.add_argument("--t_keep", type=float, default=0.5) - - parser.add_argument("--assumed_I0", type=int, default=None) - - args = parser.parse_args() - args.device = resolve_device(args.device) - return args - - -def main(): - args = get_args() - graph_sizes = parse_int_list(args.graph_sizes) - if len(graph_sizes) == 0: - raise ValueError("--graph_sizes is empty.") - - os.makedirs(args.output_dir, exist_ok=True) - - csv_path = os.path.join(args.output_dir, "runtime_breakdown_ba_sparse_rhs.csv") - fig_path = os.path.join(args.output_dir, "runtime_breakdown_ba_sparse_rhs.png") - - print("=" * 80, flush=True) - print("BA-SIR runtime breakdown experiment (sparse RHS patch)", flush=True) - print(f"device : {args.device}", flush=True) - print(f"graph_sizes : {graph_sizes}", flush=True) - print(f"T : {args.T}", flush=True) - print(f"obs_time : {parse_obs_time(args.obs_time, args.T)}", flush=True) - print(f"q_steps : {args.q_steps}", flush=True) - print("=" * 80, flush=True) - - rows = [] - total_runs = len(graph_sizes) - - for run_order, requested_graph_size in enumerate(graph_sizes, start=1): - print(f"\n=== [{run_order}/{total_runs}] running BA-SIR with n={requested_graph_size} ===", flush=True) - row = run_one_size( - requested_graph_size=requested_graph_size, - run_order=run_order, - args=args, - ) - rows.append(row) - save_rows(rows, csv_path) - - if row["status"] == "ok": - print( - "[ok] " - f"n={row['graph_size_actual']:,}, " - f"m={row['num_edges']:,}, " - f"build={row['graph_build_sec']:.2f}s, " - f"beta={row['beta_est_sec']:.2f}s, " - f"q_train={row['proposal_train_sec']:.2f}s, " - f"mcmc={row['mcmc_sec']:.2f}s, " - f"compose={row['compose_sec']:.2f}s, " - f"total_algo={row['total_algo_sec']:.2f}s, " - f"total_wall={row['total_wall_sec']:.2f}s, " - f"est_pI={row['est_pI']:.4f}, " - f"est_pR={row['est_pR']:.4f}", - flush=True, - ) - else: - print( - "[failed] " - f"n={requested_graph_size:,}, " - f"stage={row['fail_stage']}, " - f"error={row['error']}", - flush=True, - ) - - make_plot(rows, fig_path, q_steps=args.q_steps) - - print("\nSaved:") - print(f" CSV : {csv_path}", flush=True) - print(f" FIG : {fig_path}", flush=True) - - -if __name__ == "__main__": - main() \ No newline at end of file