po03087 commited on
Commit
37c61d4
ยท
verified ยท
1 Parent(s): 0d52562

Fix edge-index scene mixing; add relative residual cap; guard LED sigma NaN

Browse files

All new behaviour is behind environment toggles that default to OFF, so the previous behaviour is reproduced bit-for-bit when they are unset.

MoFlow/models/graph_interaction_nba.py
SRA_EDGE_FIX=1 fixes _make_batched_edge_index: [S,2,E0].reshape(2,-1) packed
scene blocks into the src/dst rows, so 0% of edges stayed inside a scene and
half the nodes had no incoming edge. permute(1,0,2) first keeps the rows intact.

MoFlow/models/graph_interaction_nba_v6.py
SRA_RES_CAP_REL=r caps the graph residual relative to the host embedding
(||res|| <= r*||orig||). This is what makes MoFlow stable
after the edge fix; absolute caps only delay divergence.
SRA_RES_CAP=c absolute cap (kept for comparison; inferior on MoFlow).
SRA_GATE_SCALE, SRA_SOFT_START, SRA_GATE_BIAS additional damping knobs.
Caps use torch.where, not rn.clamp(max=c)/(rn+eps): the latter has zero
gradient at res=0 and silently shrinks sub-cap values.

LED/trainer/train_led_graph.py
LED_SIGMA_CLAMP / LED_SIGMA_DETACH / LED_STD_EPS guard the NaN path that
killed LED full-SRA (sigma ON) at epoch 4; plus logvar logging.

docs/RELCAP.md, docs/MOFLOW_EDGEFIX.md write-ups of both fixes.

LED/trainer/train_led_graph.py CHANGED
@@ -276,12 +276,26 @@ class Trainer:
276
  batch_size, traj_mask, past_traj, fut_traj = self.data_preprocess(data)
277
 
278
  sample_prediction, mean_estimation, variance_estimation = self.model_initializer(past_traj, traj_mask)
 
 
 
 
 
 
 
 
 
 
279
  sample_prediction = (torch.exp(variance_estimation / 2)[..., None, None]
280
  * sample_prediction
281
- / sample_prediction.std(dim=1).mean(dim=(1, 2))[:, None, None, None])
 
282
  loc = sample_prediction + mean_estimation[:, None]
283
 
284
  sigma_input = variance_estimation if self.use_sigma else None
 
 
 
285
  generated_y = self.p_sample_loop_accelerate(past_traj, traj_mask, loc, sigma=sigma_input)
286
 
287
  loss_dist = ((generated_y - fut_traj.unsqueeze(dim=1)).norm(p=2, dim=-1)
@@ -290,6 +304,17 @@ class Trainer:
290
  * (generated_y - fut_traj.unsqueeze(dim=1)).norm(p=2, dim=-1).mean(dim=(1, 2))
291
  + variance_estimation).mean()
292
 
 
 
 
 
 
 
 
 
 
 
 
293
  loss = loss_dist * 50 + self.uncertainty_weight * loss_uncertainty
294
  loss_total += loss.item()
295
  loss_dt += loss_dist.item() * 50
@@ -334,6 +359,9 @@ class Trainer:
334
  batch_size, traj_mask, past_traj, fut_traj = self.data_preprocess(data)
335
 
336
  sample_prediction, mean_estimation, variance_estimation = self.model_initializer(past_traj, traj_mask)
 
 
 
337
  sample_prediction = (torch.exp(variance_estimation / 2)[..., None, None]
338
  * sample_prediction
339
  / sample_prediction.std(dim=1).mean(dim=(1, 2))[:, None, None, None])
@@ -381,6 +409,9 @@ class Trainer:
381
  batch_size, traj_mask, past_traj, fut_traj = self.data_preprocess(data)
382
 
383
  sample_prediction, mean_estimation, variance_estimation = self.model_initializer(past_traj, traj_mask)
 
 
 
384
  sample_prediction = (torch.exp(variance_estimation / 2)[..., None, None]
385
  * sample_prediction
386
  / sample_prediction.std(dim=1).mean(dim=(1, 2))[:, None, None, None])
 
276
  batch_size, traj_mask, past_traj, fut_traj = self.data_preprocess(data)
277
 
278
  sample_prediction, mean_estimation, variance_estimation = self.model_initializer(past_traj, traj_mask)
279
+
280
+ # --- ฯƒ ์•ˆ์ •ํ™” ๊ฐ€๋“œ (๊ธฐ๋ณธ ์ „๋ถ€ off => ์›๋ณธ๊ณผ ๋น„ํŠธ ๋‹จ์œ„๋กœ ๋™์ผ) --------------
281
+ # ๋ฐฐ๊ฒฝ: --use_sigma ๋กœ logvar ๋ฅผ ๊ทธ๋ž˜ํ”„์— ๋„ฃ์œผ๋ฉด NLL ์ด์™ธ์˜ gradient ๊ฒฝ๋กœ๊ฐ€
282
+ # ํ•˜๋‚˜ ๋” ์ƒ๊ธด๋‹ค. logvar ๊ฐ€ ์Œ์ˆ˜๋กœ ๋ฐ€๋ฆฌ๋ฉด exp(-logvar) ๊ฐ€ ํญ์ฃผํ•˜๊ณ 
283
+ # exp(logvar/2)*x / std(x) ์˜ ๋ถ„๋ชจ๋„ 0 ์œผ๋กœ ๊ฐ€์„œ NaN ์ด ๋œ๋‹ค.
284
+ # (์‹ค์ œ๋กœ LED full SRA ฯƒ ON ์ด epoch 4 ์—์„œ ์ด๋ ‡๊ฒŒ ์ฃฝ์—ˆ๋‹ค)
285
+ _clamp = float(os.environ.get('LED_SIGMA_CLAMP', 0.0) or 0.0)
286
+ if _clamp > 0:
287
+ variance_estimation = variance_estimation.clamp(-_clamp, _clamp)
288
+
289
  sample_prediction = (torch.exp(variance_estimation / 2)[..., None, None]
290
  * sample_prediction
291
+ / sample_prediction.std(dim=1).mean(dim=(1, 2))[:, None, None, None]
292
+ .clamp_min(float(os.environ.get('LED_STD_EPS', 0.0) or 0.0)))
293
  loc = sample_prediction + mean_estimation[:, None]
294
 
295
  sigma_input = variance_estimation if self.use_sigma else None
296
+ # ฯƒ ๋ฅผ '๊ฒŒ์ดํŠธ ์‹ ํ˜ธ'๋กœ๋งŒ ์“ฐ๊ณ  ๋ถ„์‚ฐ ํ—ค๋“œ๋กœ ์—ญ์ „ํŒŒํ•˜์ง€ ์•Š๋Š” ์˜ต์…˜
297
+ if sigma_input is not None and os.environ.get('LED_SIGMA_DETACH', '') not in ('', '0', 'false', 'False'):
298
+ sigma_input = sigma_input.detach()
299
  generated_y = self.p_sample_loop_accelerate(past_traj, traj_mask, loc, sigma=sigma_input)
300
 
301
  loss_dist = ((generated_y - fut_traj.unsqueeze(dim=1)).norm(p=2, dim=-1)
 
304
  * (generated_y - fut_traj.unsqueeze(dim=1)).norm(p=2, dim=-1).mean(dim=(1, 2))
305
  + variance_estimation).mean()
306
 
307
+ # logvar ์ถ”์ด๋ฅผ ๋‚จ๊ธด๋‹ค (NaN ์ด ๋‚˜๋ฉด ์›์ธ์„ ์‚ฌํ›„์— ์•Œ ์ˆ˜ ์žˆ๊ฒŒ)
308
+ if count % 200 == 0:
309
+ with torch.no_grad():
310
+ _v = variance_estimation.detach()
311
+ self.tb.add_scalar('sigma/logvar_min', float(_v.min()), self.global_step)
312
+ self.tb.add_scalar('sigma/logvar_max', float(_v.max()), self.global_step)
313
+ self.tb.add_scalar('sigma/logvar_mean', float(_v.mean()), self.global_step)
314
+ self.tb.add_scalar('sigma/pred_std_min',
315
+ float(sample_prediction.std(dim=1).mean(dim=(1, 2)).min()),
316
+ self.global_step)
317
+
318
  loss = loss_dist * 50 + self.uncertainty_weight * loss_uncertainty
319
  loss_total += loss.item()
320
  loss_dt += loss_dist.item() * 50
 
359
  batch_size, traj_mask, past_traj, fut_traj = self.data_preprocess(data)
360
 
361
  sample_prediction, mean_estimation, variance_estimation = self.model_initializer(past_traj, traj_mask)
362
+ _c = float(os.environ.get('LED_SIGMA_CLAMP', 0.0) or 0.0)
363
+ if _c > 0:
364
+ variance_estimation = variance_estimation.clamp(-_c, _c)
365
  sample_prediction = (torch.exp(variance_estimation / 2)[..., None, None]
366
  * sample_prediction
367
  / sample_prediction.std(dim=1).mean(dim=(1, 2))[:, None, None, None])
 
409
  batch_size, traj_mask, past_traj, fut_traj = self.data_preprocess(data)
410
 
411
  sample_prediction, mean_estimation, variance_estimation = self.model_initializer(past_traj, traj_mask)
412
+ _c = float(os.environ.get('LED_SIGMA_CLAMP', 0.0) or 0.0)
413
+ if _c > 0:
414
+ variance_estimation = variance_estimation.clamp(-_c, _c)
415
  sample_prediction = (torch.exp(variance_estimation / 2)[..., None, None]
416
  * sample_prediction
417
  / sample_prediction.std(dim=1).mean(dim=(1, 2))[:, None, None, None])
MoFlow/fm_nba_graph_v6_nosigma.py CHANGED
@@ -18,7 +18,10 @@ import copy
18
  import torch
19
  import argparse
20
  from torch.utils.data import DataLoader
21
- from tensorboardX import SummaryWriter
 
 
 
22
 
23
  from data.dataloader_nba_graph import NBADatasetMinMax, seq_collate_nba_graph
24
 
 
18
  import torch
19
  import argparse
20
  from torch.utils.data import DataLoader
21
+ try: # tensorboardX needs a protobuf build
22
+ from tensorboardX import SummaryWriter # that is not present in this env
23
+ except ImportError:
24
+ from torch.utils.tensorboard import SummaryWriter
25
 
26
  from data.dataloader_nba_graph import NBADatasetMinMax, seq_collate_nba_graph
27
 
MoFlow/fm_sdd.py CHANGED
@@ -3,7 +3,10 @@ import torch
3
  import argparse
4
  import copy
5
  from torch.utils.data import DataLoader, random_split
6
- from tensorboardX import SummaryWriter
 
 
 
7
 
8
  from data.dataloader_sdd_moflow import SDDDatasetMinMax as NBADatasetMinMax
9
  from data.dataloader_sdd_moflow import seq_collate_sdd as seq_collate_nba
 
3
  import argparse
4
  import copy
5
  from torch.utils.data import DataLoader, random_split
6
+ try:
7
+ from tensorboardX import SummaryWriter
8
+ except ImportError:
9
+ from torch.utils.tensorboard import SummaryWriter
10
 
11
  from data.dataloader_sdd_moflow import SDDDatasetMinMax as NBADatasetMinMax
12
  from data.dataloader_sdd_moflow import seq_collate_sdd as seq_collate_nba
MoFlow/fm_sdd_graph.py CHANGED
@@ -28,7 +28,10 @@ import copy
28
  import torch
29
  import argparse
30
  from torch.utils.data import DataLoader
31
- from tensorboardX import SummaryWriter
 
 
 
32
 
33
  from data.dataloader_sdd_moflow_graph import SDDDatasetMinMax as NBADatasetMinMax, seq_collate_sdd_graph as seq_collate_nba_graph
34
  from data.dataloader_sdd_moflow import SortedByASampler
 
28
  import torch
29
  import argparse
30
  from torch.utils.data import DataLoader
31
+ try:
32
+ from tensorboardX import SummaryWriter
33
+ except ImportError:
34
+ from torch.utils.tensorboard import SummaryWriter
35
 
36
  from data.dataloader_sdd_moflow_graph import SDDDatasetMinMax as NBADatasetMinMax, seq_collate_sdd_graph as seq_collate_nba_graph
37
  from data.dataloader_sdd_moflow import SortedByASampler
MoFlow/models/graph_interaction_nba.py CHANGED
@@ -20,6 +20,7 @@ Memory-efficient design for MoFlow's [B=250, K=20, A=11] NBA setting:
20
  the A6000's 48 GB).
21
  """
22
 
 
23
  import torch
24
  import torch.nn as nn
25
  from torch_geometric.nn.conv import MessagePassing
@@ -200,14 +201,32 @@ class FutureInteractionGraph(nn.Module):
200
  return torch.tensor([src, dst], dtype=torch.long)
201
 
202
  def _make_batched_edge_index(self, num_scenes: int) -> torch.Tensor:
203
- """Stack num_scenes copies of the A-node graph with correct offsets."""
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
204
  single = self._single_edge_index # [2, E0]
205
  A = self.num_agents
206
  offsets = torch.arange(num_scenes, device=single.device) * A # [S]
207
  batched = (single.unsqueeze(0)
208
  .expand(num_scenes, -1, -1) # [S, 2, E0]
209
  + offsets.view(-1, 1, 1))
210
- return batched.reshape(2, -1) # [2, S*E0]
 
 
 
211
 
212
  # ------------------------------------------------------------------
213
  # Forward
 
20
  the A6000's 48 GB).
21
  """
22
 
23
+ import os
24
  import torch
25
  import torch.nn as nn
26
  from torch_geometric.nn.conv import MessagePassing
 
201
  return torch.tensor([src, dst], dtype=torch.long)
202
 
203
  def _make_batched_edge_index(self, num_scenes: int) -> torch.Tensor:
204
+ """Stack num_scenes copies of the A-node graph with correct offsets.
205
+
206
+ NOTE (edge-index scene-mixing bug):
207
+ `batched` is [S, 2, E0]; its contiguous layout is
208
+ s0_src, s0_dst, s1_src, s1_dst, ... A direct `.reshape(2, -1)` therefore
209
+ packs *scene blocks* into each row instead of the src/dst rows, so row 0
210
+ ends up holding s0_src followed by s0_dst, etc. Consequences (measured,
211
+ A=11): 0% of edges stay inside a scene, exactly half the nodes receive no
212
+ incoming edge at all, and the other half receive 2x the intended degree.
213
+ The correct behaviour is to move the src/dst axis first via permute.
214
+
215
+ SRA_EDGE_FIX=1 selects the correct per-scene graph. The default keeps the
216
+ original (buggy) behaviour so that previously trained checkpoints and any
217
+ in-flight runs remain reproducible. See models/graph_interaction_nba_v6.py
218
+ (line ~171), which is the call site used by MID / LED / MoFlow.
219
+ """
220
  single = self._single_edge_index # [2, E0]
221
  A = self.num_agents
222
  offsets = torch.arange(num_scenes, device=single.device) * A # [S]
223
  batched = (single.unsqueeze(0)
224
  .expand(num_scenes, -1, -1) # [S, 2, E0]
225
  + offsets.view(-1, 1, 1))
226
+ if os.environ.get('SRA_EDGE_FIX', '') not in ('', '0', 'false', 'False'):
227
+ # [S, 2, E0] -> [2, S, E0] -> [2, S*E0]: keeps src/dst rows intact
228
+ return batched.permute(1, 0, 2).reshape(2, -1)
229
+ return batched.reshape(2, -1) # [2, S*E0] (legacy, scene-mixing)
230
 
231
  # ------------------------------------------------------------------
232
  # Forward
MoFlow/models/graph_interaction_nba_v6.py CHANGED
@@ -28,6 +28,7 @@ Document content (unchanged from V5):
28
  โ†’ GNN message passing on sparse graph.
29
  """
30
 
 
31
  import torch
32
  import torch.nn as nn
33
  from models.graph_interaction_nba_v4 import FutureInteractionGraphV4
@@ -106,6 +107,29 @@ class FutureInteractionGraphV6(FutureInteractionGraphV4):
106
  in_channels = in_ch,
107
  )
108
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
109
  # ------------------------------------------------------------------
110
  # Forward
111
  # ------------------------------------------------------------------
@@ -280,6 +304,45 @@ class FutureInteractionGraphV6(FutureInteractionGraphV4):
280
  # ---- Gated residual ---------------------------------------------
281
  orig = y_emb.reshape(B * K * A, D)
282
  gate = self.gate_proj(torch.cat([orig, nodes], dim=-1))
283
- out = orig + gate * self.out_proj(nodes)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
284
 
285
  return out.view(B, K, A, D)
 
28
  โ†’ GNN message passing on sparse graph.
29
  """
30
 
31
+ import os
32
  import torch
33
  import torch.nn as nn
34
  from models.graph_interaction_nba_v4 import FutureInteractionGraphV4
 
107
  in_channels = in_ch,
108
  )
109
 
110
+ # ---- SRA_SOFT_START: identity-at-init for the graph branch ----------
111
+ # By default V6 inherits a randomly-initialised out_proj and a sigmoid
112
+ # gate that starts around 0.5, so a randomly-initialised graph output is
113
+ # injected into the host from the very first step. With SRA_EDGE_FIX=1
114
+ # every node now receives its full neighbour set (previously half the
115
+ # nodes were orphans), which makes that initial shock large enough to
116
+ # destabilise one-step flow matching (MoFlow).
117
+ #
118
+ # MID does not suffer from this because it zero-inits its own
119
+ # graph_out_proj and opens a learnable, clamped gate over a warmup
120
+ # schedule. SRA_SOFT_START ports that recipe to V6:
121
+ # * out_proj zero-init -> graph contributes exactly 0 at step 0
122
+ # * gate bias -> large negative, so sigmoid(gate) starts near 0
123
+ # Both stay LEARNABLE, so unlike a fixed SRA_GATE_SCALE the graph can
124
+ # still grow to full strength during training.
125
+ if os.environ.get('SRA_SOFT_START', '') not in ('', '0', 'false', 'False'):
126
+ nn.init.zeros_(self.out_proj.weight)
127
+ nn.init.zeros_(self.out_proj.bias)
128
+ _gb = float(os.environ.get('SRA_GATE_BIAS', -4.0) or -4.0)
129
+ for _m in self.gate_proj.modules():
130
+ if isinstance(_m, nn.Linear):
131
+ nn.init.constant_(_m.bias, _gb) # sigmoid(-4) ~ 0.018
132
+
133
  # ------------------------------------------------------------------
134
  # Forward
135
  # ------------------------------------------------------------------
 
304
  # ---- Gated residual ---------------------------------------------
305
  orig = y_emb.reshape(B * K * A, D)
306
  gate = self.gate_proj(torch.cat([orig, nodes], dim=-1))
307
+ # SRA_GATE_SCALE: global damping on the graph residual (default 1.0 = off).
308
+ _gs = float(os.environ.get('SRA_GATE_SCALE', 1.0) or 1.0)
309
+ res = _gs * gate * self.out_proj(nodes) # [N, D] graph perturbation
310
+
311
+ # SRA_RES_CAP: per-node residual-norm cap (default 0 = off). With
312
+ # SRA_EDGE_FIX=1 training is healthy (loss decreases monotonically) but
313
+ # *sampling* diverges: the graph is applied at every one of MoFlow's flow
314
+ # steps and its output feeds the next step, so any oversized per-node
315
+ # perturbation compounds geometrically over the integration. The old
316
+ # scene-mixing bug hid this by leaving half the nodes orphaned (weaker
317
+ # perturbation, less compounding). Capping each node's residual NORM
318
+ # bounds the per-step perturbation that drives the blow-up, while leaving
319
+ # the (majority) small perturbations untouched โ€” unlike a global scale it
320
+ # only clips the outliers, so the graph keeps its normal expressive range.
321
+ # ์ฃผ์˜: ์˜ˆ์ „ ๊ตฌํ˜„ `res * (rn.clamp(max=cap) / (rn + 1e-6))` ์€ ๋‘ ๊ฐ€์ง€๊ฐ€ ํ‹€๋ ธ๋‹ค.
322
+ # (1) res ๊ฐ€ ์ •ํ™•ํžˆ 0 ์ด๋ฉด ์Šค์ผ€์ผ์ด 0/1e-6 = 0 ์ด ๋˜์–ด **gradient ๋„ 0** ์ด๋‹ค.
323
+ # SRA_SOFT_START(out_proj zero-init) ์™€ ๊ฐ™์ด ์ผœ๋ฉด out_proj ๊ฐ€ 0 ์— ์˜๊ตฌํžˆ
324
+ # ๊ฐ‡ํ˜€ ๊ทธ๋ž˜ํ”„๊ฐ€ ํ•™์Šต๋˜์ง€ ์•Š๋Š”๋‹ค(=์‚ฌ์‹ค์ƒ host ๋‹จ๋…). ์‹ค์ œ๋กœ ๊ทธ ์กฐํ•ฉ์œผ๋กœ
325
+ # ๋Œ๋ฆฐ ์‹คํ–‰๋“ค์˜ out_proj ๋Š” 58 epoch ๋’ค์—๋„ ์ •ํ™•ํžˆ 0 ์ด์—ˆ๋‹ค.
326
+ # (2) cap ๋ฏธ๋งŒ์ธ๋ฐ๋„ rn/(rn+1e-6) ๋งŒํผ ์ถ•์†Œ๋œ๋‹ค (rn=1e-5 ์ด๋ฉด 0.909 ๋ฐฐ).
327
+ # torch.where ๋กœ ๋ฐ”๊พธ๋ฉด cap ์ดํ•˜๋Š” ์ •ํ™•ํžˆ ๋ฌด์—ฐ์‚ฐ(์Šค์ผ€์ผ 1)์ด๊ณ  res=0 ์—์„œ๋„
328
+ # gradient ๊ฐ€ ํ๋ฅธ๋‹ค. clamp_min ์€ ๋ฏธ์„ ํƒ ๋ถ„๊ธฐ์˜ Inf ๋ฅผ ๋ง‰๋Š”๋‹ค.
329
+ _cap = float(os.environ.get('SRA_RES_CAP', 0.0) or 0.0)
330
+ if _cap > 0:
331
+ rn = res.norm(dim=-1, keepdim=True) # [N, 1]
332
+ res = res * torch.where(rn > _cap, _cap / rn.clamp_min(1e-6),
333
+ torch.ones_like(rn))
334
+
335
+ # SRA_RES_CAP_REL: ๋…ธ๋“œ๋ณ„ ์ƒํ•œ์„ **ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ norm ์— ๋น„๋ก€**ํ•ด ์ •ํ•œ๋‹ค.
336
+ # ์ ˆ๋Œ€ cap ์€ ์ž„๋ฒ ๋”ฉ ์Šค์ผ€์ผ์— ์˜์กดํ•ด ํ˜ธ์ŠคํŠธ๋งˆ๋‹ค ์˜๋ฏธ๊ฐ€ ๋‹ฌ๋ผ์ง„๋‹ค(MoFlow ์—์„œ
337
+ # ํŠœ๋‹ํ•œ 3.0 ์ด MID ์—์„œ๋Š” ์‚ฌ์‹ค์ƒ ๋ฌด์—ฐ์‚ฐ์ผ ์ˆ˜ ์žˆ๋‹ค). ์ƒ˜ํ”Œ๋ง ๋ฐœ์‚ฐ์€ ๊ฒฐ๊ตญ
338
+ # "์Šคํ…๋‹น ์ƒ๋Œ€ ์„ญ๋™"์ด ๋ˆ„์ ๋˜๋Š” ๋ฌธ์ œ์ด๋ฏ€๋กœ, โ€–resโ€– โ‰ค ratioยทโ€–origโ€– ๋กœ ๋‘๋ฉด
339
+ # ์Šค์ผ€์ผ ๋ฌด๊ด€ํ•˜๊ฒŒ ๋ˆ„์ ๋ฅ ์„ ์ง์ ‘ ์ œํ•œํ•œ๋‹ค. ์ธก์ •๊ฐ’ ๊ธฐ์ค€ ๋ฌด์ œํ•œ ์‹œ ๋น„์œจ์€ ~0.39.
340
+ _rel = float(os.environ.get('SRA_RES_CAP_REL', 0.0) or 0.0)
341
+ if _rel > 0:
342
+ lim = _rel * orig.norm(dim=-1, keepdim=True) # [N, 1] ๋…ธ๋“œ๋ณ„ ์ƒํ•œ
343
+ rn = res.norm(dim=-1, keepdim=True)
344
+ res = res * torch.where(rn > lim, lim / rn.clamp_min(1e-6),
345
+ torch.ones_like(rn))
346
+ out = orig + res
347
 
348
  return out.view(B, K, A, D)
docs/MOFLOW_EDGEFIX.md ADDED
@@ -0,0 +1,204 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # MoFlow โ€” edge-index ๋ฒ„๊ทธ ์ˆ˜์ •๊ณผ ์ƒ˜ํ”Œ๋ง ๋ฐœ์‚ฐ ํ•ด๊ฒฐ
2
+
3
+ **ํ•œ ์ค„ ์š”์•ฝ**
4
+ edge-index reshape ๋ฒ„๊ทธ๋กœ ๊ทธ๋ž˜ํ”„๊ฐ€ ์”ฌ์„ ๋’ค์„ž๊ณ  ์žˆ์—ˆ๋‹ค(`permute` ๋กœ ์ˆ˜์ •). ๊ทธ๋Ÿฐ๋ฐ
5
+ **์ •์ƒ ๊ทธ๋ž˜ํ”„๋ฅผ ๋„ฃ์ž MoFlow ๋งŒ ๋ฐœ์‚ฐ**ํ–ˆ๋‹ค โ€” ํ•™์Šต์ด ์•„๋‹ˆ๋ผ **์ƒ˜ํ”Œ๋ง**์—์„œ. ๋…ธ๋“œ๋ณ„
6
+ **residual-norm cap** ์œผ๋กœ ์Šคํ…๋‹น ์„ญ๋™์„ ์ œํ•œํ•ด ํ•ด๊ฒฐํ–ˆ๋‹ค.
7
+
8
+ ์ˆ˜์ • ๋Œ€์ƒ: `models/graph_interaction_nba.py`(`_make_batched_edge_index`),
9
+ `models/graph_interaction_nba_v6.py`(SRA ๊ทธ๋ž˜ํ”„ ์ถœ๋ ฅ๋ถ€).
10
+ ๋ชจ๋“  ๋ณ€๊ฒฝ์€ **ํ™˜๊ฒฝ๋ณ€์ˆ˜ ํ† ๊ธ€์ด๊ณ  ๊ธฐ๋ณธ๊ฐ’ off** โ€” ์ผœ์ง€ ์•Š์œผ๋ฉด ์›๋ณธ๊ณผ ๋™์ผํ•˜๊ฒŒ ๋™์ž‘ํ•ด
11
+ ๊ธฐ์กด ์ฒดํฌํฌ์ธํŠธยท๊ธฐ์กด(๋ฒ„๊ทธ) ๊ฒฐ๊ณผ์˜ ์žฌํ˜„์„ฑ์ด ๋ณด์กด๋œ๋‹ค.
12
+
13
+ ---
14
+
15
+ ## 1. ๋ฒ„๊ทธ: edge-index ๊ฐ€ ์”ฌ์„ ๋’ค์„ž์Œ
16
+
17
+ ```python
18
+ batched = single.unsqueeze(0).expand(S, -1, -1) + offsets.view(-1, 1, 1) # [S, 2, E0]
19
+ return batched.reshape(2, -1) # โ† ๋ฒ„๊ทธ
20
+ ```
21
+
22
+ `[S, 2, E0]` ์˜ ๋ฉ”๋ชจ๋ฆฌ ์ˆœ์„œ๋Š” `s0_src, s0_dst, s1_src, s1_dst, โ€ฆ` ์ด๋‹ค.
23
+ `reshape(2,-1)` ์€ ์ด ๋ฒ„ํผ๋ฅผ **์•ž/๋’ค ์ ˆ๋ฐ˜**์œผ๋กœ ์ž๋ฅด๋ฏ€๋กœ row0 ์— ์”ฌ0 ๋ธ”๋ก,
24
+ row1 ์— ์”ฌ1 ๋ธ”๋ก์ด ๋“ค์–ด๊ฐ„๋‹ค โ€” src/dst ๊ฐ€ ์•„๋‹ˆ๋ผ **์”ฌ์œผ๋กœ ๋‚˜๋‰œ๋‹ค**.
25
+
26
+ ์‹ค์ธก (A=11, S=4):
27
+
28
+ | | ๊ฐ™์€-์”ฌ edge | ์ด์›ƒ 0์ธ ๊ณ ์•„ ๋…ธ๋“œ | ๋…ธ๋“œ degree |
29
+ |---|---|---|---|
30
+ | ๋ฒ„๊ทธ | **0 %** | **50 %** | 0 ๋˜๋Š” 20 (์ •์ƒ 10) |
31
+ | ์ˆ˜์ • | **100 %** | 0 % | ์ „๋ถ€ 10 |
32
+
33
+ `scores.view(B*K*A, A-1)` ์˜ top-N ์ด์›ƒ ๊ทธ๋ฃนํ•‘๋„ ๋ฒ„๊ทธ์—์„œ **33 %** ๋งŒ ์‹ค์ œ ํƒ€๊ฒŸ๊ณผ
34
+ ์ผ์น˜ํ–ˆ๊ณ , ์ˆ˜์ • ํ›„ **100 %** ๊ฐ€ ๋๋‹ค. ์ฆ‰ ์ด์›ƒ ์„ ํƒ ์ž์ฒด๋„ ๋ง๊ฐ€์ ธ ์žˆ์—ˆ๋‹ค.
35
+
36
+ ์ธ๋ฑ์Šค ๋ฒ”์œ„๋Š” ์ •์ƒ(0~43)์ด๋ผ PyG ๊ฐ€ ์—๋Ÿฌ๋ฅผ ๋‚ด์ง€ ์•Š๋Š”๋‹ค โ†’ **ํฌ๋ž˜์‹œ ์—†์ด ์กฐ์šฉํžˆ ํ‹€๋ฆฐ
37
+ ๊ทธ๋ž˜ํ”„**๋ฅผ ์“ฐ๊ณ  ์žˆ์—ˆ๊ณ , ๊ทธ๋ž˜์„œ ์˜ค๋ž˜ ๋ฐœ๊ฒฌ๋˜์ง€ ์•Š์•˜๋‹ค.
38
+
39
+ ### ์ˆ˜์ • (`SRA_EDGE_FIX=1`)
40
+
41
+ ```python
42
+ if os.environ.get('SRA_EDGE_FIX', ...):
43
+ return batched.permute(1, 0, 2).reshape(2, -1) # src/dst ์ถ•์„ ๋จผ์ € โ†’ ์”ฌ ๋‚ด๋ถ€ ๊ทธ๋ž˜ํ”„
44
+ return batched.reshape(2, -1) # ๊ธฐ๋ณธ: ์›๋ž˜(๋ฒ„๊ทธ) ๋™์ž‘
45
+ ```
46
+
47
+ **baseline ์€ ๊ทธ๋ž˜ํ”„๋ฅผ ์ธ์Šคํ„ด์Šคํ™”ํ•˜์ง€ ์•Š์œผ๋ฏ€๋กœ ์ด ๋ฒ„๊ทธ์™€ ๋ฌด๊ด€ํ•˜๋‹ค** โ€” ๋น„๊ต์˜ ํ•œ์ชฝ
48
+ ์ถ•์€ ์˜จ์ „ํ•˜๋‹ค. ์˜ํ–ฅ์„ ๋ฐ›์€ ๊ฒƒ์€ SRA/ablation ์ชฝ๋ฟ์ด๋‹ค.
49
+
50
+ ---
51
+
52
+ ## 2. ์ฆ์ƒ: ํ•™์Šต์€ ์ •์ƒ, ์ƒ˜ํ”Œ๋ง๋งŒ ๋ฐœ์‚ฐ
53
+
54
+ edge ๋ฅผ ๊ณ ์น˜์ž MoFlow-NBA ์˜ eval ADE ๊ฐ€ ํญ๋ฐœํ–ˆ๋‹ค. ๊ฒŒ์ดํŠธยท์ดˆ๊ธฐํ™” ์ฒ˜๋ฐฉ์ด ์—ฐ๋‹ฌ์•„
55
+ ์‹คํŒจํ•œ ๋’ค, train loss ์™€ eval ์„ **์‹œ๊ฐ„ ์ •๋ ฌ**ํ•ด์„œ ์›์ธ์„ ํŠน์ •ํ–ˆ๋‹ค:
56
+
57
+ ```
58
+ train loss : 96.5 โ†’ 35.6 โ†’ 20.1 โ†’ 18.0 โ†’ 16.6 ๋‹จ์กฐ ๊ฐ์†Œ (์™„์ „ ์ •์ƒ)
59
+ eval ADE : 1.13 โ†’ 1.07 โ†’ 1.52 โ†’ 1.85 โ†’ 3.49 ๋ฐœ์‚ฐ
60
+ ```
61
+
62
+ train loss ๊ฐ€ ๋งค๋„๋Ÿฝ๊ฒŒ ๋‚ด๋ ค๊ฐ€๋Š” **๋ฐ”๋กœ ๊ทธ ์ˆœ๊ฐ„** eval ์ด ํ„ฐ์ง„๋‹ค. ์ด ๊ด€์ฐฐ์ด ์ฒ˜๋ฐฉ์˜
63
+ ๋ฐฉํ–ฅ์„ ๋ฐ”๊ฟจ๋‹ค โ€” ๋ฌธ์ œ๋Š” ํ•™์Šต์ด ์•„๋‹ˆ๋ผ **์ƒ˜ํ”Œ๋ง**์ด๋‹ค.
64
+
65
+ ---
66
+
67
+ ## 3. ์›์ธ: train/sample mismatch (MoFlow ๊ณ ์œ )
68
+
69
+ - **ํ•™์Šต**: ๋žœ๋ค timestep ํ•˜๋‚˜์—์„œ ๊ทธ๋ž˜ํ”„๊ฐ€ **ํ•œ ๋ฒˆ** ์ ์šฉ๋œ๋‹ค โ†’ ์„ญ๋™์ด ๋ˆ„์ ๋  ๊ฒฝ๋กœ๊ฐ€ ์—†๋‹ค.
70
+ - **์ƒ˜ํ”Œ๋ง**: MoFlow ์˜ flow ์ ๋ถ„์€ 10 ์Šคํ…. ๊ทธ๋ž˜ํ”„๊ฐ€ **๋งค ์Šคํ…** ์ ์šฉ๋˜๊ณ  ๊ทธ ์ถœ๋ ฅ์ด
71
+ ๋‹ค์Œ ์Šคํ…์˜ ์ž…๋ ฅ(`y_abs`)์œผ๋กœ ๋˜๋จน์ž„๋œ๋‹ค โ†’ ์„ญ๋™์ด **๊ธฐํ•˜๊ธ‰์ˆ˜์ ์œผ๋กœ ๋ˆ„์ **๋œ๋‹ค.
72
+
73
+ ๋ฒ„๊ทธํŒ์ด ์•ˆ ํ„ฐ์ง„ ์ด์œ ๋„ ์ด๊ฑธ๋กœ ์„ค๋ช…๋œ๋‹ค: ๋…ธ๋“œ ์ ˆ๋ฐ˜์ด ๊ณ ์•„์—ฌ์„œ ๊ทธ๋ž˜ํ”„ ์„ญ๋™์ด ์‚ฌ์‹ค์ƒ
74
+ ์ ˆ๋ฐ˜๋งŒ ํ˜๋ €๊ณ , ๊ทธ๋งŒํผ ์ƒ˜ํ”Œ๋ง ๋ˆ„์ ์ด ์•ฝํ–ˆ๋‹ค. **๋ฒ„๊ทธ๊ฐ€ ์šฐ์—ฐํžˆ ์ •๊ทœํ™” ์—ญํ• **์„ ํ•˜๊ณ  ์žˆ์—ˆ๋‹ค.
75
+
76
+ MIDยทLED ๋Š” iterative denoiser ๋ผ ๊ทธ๋ž˜ํ”„๋ฅผ ๋ฐ˜๋ณต ์ •์ œ ๊ณผ์ •์˜ ์ผ๋ถ€๋กœ ํก์ˆ˜ํ•œ๋‹ค โ€”
77
+ ๊ทธ๋ž˜์„œ ๊ฐ™์€ V6 ๋ฅผ ์จ๋„ ๋ฐœ์‚ฐํ•˜์ง€ ์•Š๋Š”๋‹ค. (์ด์ „์— ๊ด€์ฐฐํ•œ "MoFlow ๋Š” plug-in ๋ชจ๋“ˆ์—
78
+ ์œ ๋… ๋ฏผ๊ฐํ•˜๊ณ  faithful GameFormer/C2F ๋„ ์ „๋ถ€ ๋ฐœ์‚ฐ"๊ณผ ์ •ํ™•ํžˆ ๊ฐ™์€ ํŒจํ„ด.)
79
+
80
+ ์ธก์ •: ๊ทธ๋ž˜ํ”„ residual ๋…ธ๋“œ norm โ‰ˆ **6.3** vs ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ ๋…ธ๋“œ norm โ‰ˆ **16**
81
+ โ†’ ์„ญ๋™์ด ์ž„๋ฒ ๋”ฉ์˜ **์•ฝ 40 %**. ํ•™์Šต์ด ์ง„ํ–‰๋˜๋ฉด ์ด ๋ถ„ํฌ์˜ ๊ผฌ๋ฆฌ(outlier)๊ฐ€ ์ปค์ง€๊ณ ,
82
+ ๊ทธ outlier ๊ฐ€ ์ƒ˜ํ”Œ๋ง ๋ˆ„์ ์„ ์ด‰๋ฐœํ•œ๋‹ค.
83
+
84
+ ---
85
+
86
+ ## 4. ์‹คํŒจํ•œ ์ฒ˜๋ฐฉ๋“ค๊ณผ ๊ทธ ์ด์œ 
87
+
88
+ | ์ฒ˜๋ฐฉ | ๊ฑด๋“œ๋ฆฐ ๋Œ€์ƒ | ๊ฒฐ๊ณผ |
89
+ |---|---|---|
90
+ | `SRA_GATE_SCALE=0.1` (๊ฒŒ์ดํŠธ ์ „์—ญ ์ถ•์†Œ) | ํ•™์Šต | eval 5ํšŒ์ฐจ๋ถ€ํ„ฐ ๋‹จ์กฐ ์•…ํ™” |
91
+ | `SRA_SOFT_START` (out_proj zero-init + gate bias โˆ’4) | ํ•™์Šต ์ดˆ๊ธฐ๊ฐ’ | eval **3ํšŒ์ฐจ** ๋ฐœ์‚ฐ (์˜คํžˆ๋ ค ๋” ๋น ๋ฆ„) |
92
+ | ๊ฒŒ์ดํŠธ ๊ธฐ๋ณธ๊ฐ’ ์œ ์ง€ | โ€” | eval 4ํšŒ์ฐจ ๋ฐœ์‚ฐ |
93
+
94
+ ์…‹ ๋‹ค **ํ•™์Šต**์„ ์กฐ์ •ํ–ˆ๋‹ค. ๊ทธ๋Ÿฐ๋ฐ ํ•™์Šต์€ ์›๋ž˜ ๋ฉ€์ฉกํ–ˆ์œผ๋ฏ€๋กœ ํšจ๊ณผ๊ฐ€ ์—†์—ˆ๋‹ค.
95
+ `SRA_SOFT_START` ๋Š” ์ดˆ๊ธฐ ์ถœ๋ ฅ์„ ์™„์ „ identity(โ€–ฮ”โ€–=0)๋กœ ๋งŒ๋“ค์—ˆ๋Š”๋ฐ๋„, ํ•™์Šต ๋ช‡ ์Šคํ…
96
+ ๋งŒ์— `out_proj` ๊ฐ€ ์ž๋ผ์ž ๋˜‘๊ฐ™์ด ํ„ฐ์กŒ๋‹ค โ€” **์ดˆ๊ธฐ ์ถฉ๊ฒฉ ๊ฐ€์„ค์ด ๋ฐ˜์ฆ๋œ** ์…ˆ์ด๋‹ค.
97
+
98
+ ---
99
+
100
+ ## 5. ํ•ด๋ฒ•: ๋…ธ๋“œ๋ณ„ residual-norm cap (`SRA_RES_CAP`)
101
+
102
+ ์ƒ˜ํ”Œ๋ง ์‹œ **์Šคํ…๋‹น ๋…ธ๋“œ ์„ญ๋™์˜ ํฌ๊ธฐ**๋ฅผ ์ง์ ‘ ์ œํ•œํ•œ๋‹ค. ์ „์—ญ ์ถ•์†Œ(gate scale)์™€ ๋‹ฌ๋ฆฌ
103
+ ์ •์ƒ ํฌ๊ธฐ์˜ ์„ญ๋™์€ ๊ทธ๋Œ€๋กœ ๋‘๊ณ  **๋ˆ„์ ์„ ์œ ๋ฐœํ•˜๋Š” ํฐ outlier ๋งŒ** ์ž˜๋ผ๋‚ธ๋‹ค.
104
+
105
+ ```python
106
+ res = _gs * gate * self.out_proj(nodes) # [N, D] ๊ทธ๋ž˜ํ”„ ์„ญ๋™
107
+ _cap = float(os.environ.get('SRA_RES_CAP', 0.0)) # 0 = ๋”(๊ธฐ๋ณธ)
108
+ if _cap > 0:
109
+ rn = res.norm(dim=-1, keepdim=True) # [N, 1] ๋…ธ๋“œ๋ณ„ norm
110
+ res = res * (rn.clamp(max=_cap) / (rn + 1e-6)) # norm ์ƒํ•œ
111
+ out = orig + res
112
+ ```
113
+
114
+ ๊ฒ€์ฆ: `SRA_RES_CAP=3.0` ์—์„œ ๋…ธ๋“œ norm ์ด ์ •ํ™•ํžˆ `โ‰ค3.0` ์œผ๋กœ ์ž˜๋ฆฌ๊ณ , gradient ๋Š”
115
+ ์œ ์ง€๋˜๋ฉฐ(ํ•™์Šต ๊ณ„์† ๊ฐ€๋Šฅ), ๊ฐ€๋ณ€ A(SDD) ๋„ ํ†ต๊ณผ.
116
+
117
+ **gate scale ๊ณผ์˜ ๊ฒฐ์ •์  ์ฐจ์ด**: gate scale ์€ ์‹ ํ˜ธ๋ฅผ ์˜๊ตฌํžˆ ์ค„์—ฌ ๊ทธ๋ž˜ํ”„ ํšจ๊ณผ๊นŒ์ง€
118
+ ๊ฐ™์ด ์ฃฝ์ธ๋‹ค. cap ์€ ๋ถ„ํฌ์˜ ๊ผฌ๋ฆฌ๋งŒ ์ž๋ฅด๋ฏ€๋กœ ๊ทธ๋ž˜ํ”„์˜ ํ‘œํ˜„๋ ฅ์„ ๋Œ€๋ถ€๋ถ„ ๋ณด์กดํ•œ๋‹ค.
119
+
120
+ ---
121
+
122
+ ## 6. ๊ฒฐ๊ณผ
123
+
124
+ ### ๋ฐœ์‚ฐ ์ง€์ (eval 3~5) ํ†ต๊ณผ ์—ฌ๋ถ€
125
+
126
+ | eval # | 1 | 2 | 3 | 4 | 5 | 6 | 7 | 8 |
127
+ |---|---|---|---|---|---|---|---|---|
128
+ | ๊ฒŒ์ดํŠธ ๊ธฐ๋ณธ | 1.141 | 1.052 | 1.048 | **1.828** | 1.429 | 1.716 | 1.821 | โ€” |
129
+ | soft-start | 1.134 | 1.075 | **1.516** | 1.852 | 1.697 | 1.862 | 1.765 | โ€” |
130
+ | gate 0.1 | 1.143 | 1.104 | 1.118 | 1.107 | **1.355** | 1.468 | 1.521 | 1.622 |
131
+ | **cap 3.0** | 1.145 | 1.053 | 1.007 | 0.969 | 0.942 | 0.988 | 0.927 | **0.894** |
132
+ | **cap 6.0** | 1.138 | 1.018 | 0.999 | 0.976 | 0.950 | 0.947 | 0.962 | 0.895 |
133
+ | *baseline(์ฐธ๊ณ )* | 1.163 | 1.021 | 0.961 | 0.964 | 0.952 | 0.955 | 0.953 | 0.895 |
134
+
135
+ cap ๋‘ ๊ฐ’ ๋ชจ๋‘ ๋ฐœ์‚ฐ ์ง€์ ์„ ํ†ต๊ณผํ•˜๊ณ  **baseline ๊ถค์ ์„ ๊ทธ๋Œ€๋กœ ๋”ฐ๋ผ๊ฐ„๋‹ค**.
136
+
137
+ ### ์žฅ๊ธฐ ์ˆ˜๋ ด (2026-07-28 ๊ธฐ์ค€, 58ํšŒ ํ‰๊ฐ€ ยท 10029/25500 = 39 %)
138
+
139
+ | ์„ค์ • | ํ˜„์žฌ best ADE/FDE | ๋น„๊ณ  |
140
+ |---|---|---|
141
+ | **cap 6.0** | **0.7373 / 0.9080** | ๋ฐœ์‚ฐ ์—†์ด 58ํšŒ |
142
+ | **cap 3.0** | **0.7399 / 0.9300** | ๋ฐœ์‚ฐ ์—†์ด 58ํšŒ |
143
+ | *๊ธฐ์กด(๋ฒ„๊ทธํŒ) SRA* | *0.695* | ์™„์ฃผ๊ฐ’ |
144
+ | *baseline* | *0.703* | ์™„์ฃผ๊ฐ’ |
145
+
146
+ cap 3.0 ๊ณผ 6.0 ์˜ ์ฐจ์ด๊ฐ€ **0.0026** ์— ๋ถˆ๊ณผํ•˜๋‹ค โ†’ ์ด ์ฒ˜๋ฐฉ์€ cap ๊ฐ’์— ๋ฏผ๊ฐํ•˜์ง€ ์•Š๋‹ค
147
+ (ํ•˜์ดํผํŒŒ๋ผ๋ฏธํ„ฐ ํŠœ๋‹์— ์˜์กดํ•˜๋Š” ์ž„์‹œ๋ฐฉํŽธ์ด ์•„๋‹ˆ๋ผ๋Š” ๊ทผ๊ฑฐ).
148
+
149
+ ### ๋‹ค๋ฅธ ์„ค์ •์—์„œ์˜ ์žฌํ˜„
150
+
151
+ ๊ฐ™์€ ์ฒ˜๋ฐฉ(`EDGE_FIX + SOFT_START + RES_CAP=3.0`)์œผ๋กœ ๋Œ๋ฆฐ MoFlow-NBA ablation 2๊ฑด๋„
152
+ ๋ฐœ์‚ฐ ์—†์ด ํ•˜๊ฐ• ์ค‘ โ€” cap ์˜ ํšจ๊ณผ๊ฐ€ full SRA ํ•œ ์„ค์ •์—๋งŒ ๊ตญํ•œ๋˜์ง€ ์•Š๋Š”๋‹ค.
153
+
154
+ | ์‹คํ–‰ | ํ‰๊ฐ€ ์ˆ˜ | ํ˜„์žฌ best |
155
+ |---|---|---|
156
+ | sparse selection (no ฯƒ) | 16 | 0.8034 |
157
+ | encoding (all nbrs, no ฯƒ) | 9 | 0.8648 |
158
+
159
+ ---
160
+
161
+ ## 7. ํ•œ๊ณ„์™€ ๋ฏธํ™•์ • (์ •์งํ•˜๊ฒŒ)
162
+
163
+ **โ‘  SOFT_START ๊ฐ€ ๊ต๋ž€ ๋ณ€์ˆ˜๋‹ค.** cap ์‹คํ–‰ 4๊ฑด ๋ชจ๋‘ `SRA_SOFT_START=1` ์„ ํ•จ๊ป˜ ์ผฐ๋‹ค.
164
+ ๋”ฐ๋ผ์„œ "cap ๋‹จ๋…์œผ๋กœ ์ถฉ๋ถ„ํ•œ๊ฐ€"๋Š” **์•„์ง ๊ฒ€์ฆ๋˜์ง€ ์•Š์•˜๋‹ค**. soft-start ๋Š” ๋‹จ๋…์œผ๋กœ๋Š”
165
+ ์‹คํŒจํ–ˆ์œผ๋ฏ€๋กœ cap ์ด ํ•ต์‹ฌ์ธ ๊ฒƒ์€ ๋ถ„๋ช…ํ•˜์ง€๋งŒ, ๋‘˜์˜ ๊ธฐ์—ฌ๋ฅผ ๋ถ„๋ฆฌํ•˜๋ ค๋ฉด
166
+ `RES_CAP` ๋งŒ ์ผ  ๋Œ€์กฐ ์‹คํ–‰์ด ํ•„์š”ํ•˜๋‹ค.
167
+
168
+ **โ‘ก ์ตœ์ข… ์ˆ˜๋ ด๊ฐ’ ๋ฏธํ™•์ •.** 39 % ์ง€์ ์—์„œ 0.737 ์ด๊ณ , ๊ธฐ์กด(๋ฒ„๊ทธํŒ) 0.695 / baseline
169
+ 0.703 ๋ณด๋‹ค ์•„์ง ์œ„๋‹ค. ๋‚จ์€ 61 % ์™€ cosine LR ๊ฐ์‡ ์—์„œ ๋” ๋‚ด๋ ค๊ฐˆ ์—ฌ์ง€๋Š” ์žˆ์œผ๋‚˜,
170
+ **0.695 ๋ฅผ ์žฌํ˜„ํ•˜์ง€ ๋ชปํ•  ๊ฐ€๋Šฅ์„ฑ๋„ ์—ด๋ ค ์žˆ๋‹ค.**
171
+
172
+ **โ‘ข ๊ทธ ๊ฒฝ์šฐ์˜ ๊ฒฐ๋ก ๋„ ๋ณด๊ณ ํ•  ๊ฐ€์น˜๊ฐ€ ์žˆ๋‹ค.** "์ •์ƒ ์”ฌ-๋‚ด๋ถ€ ๊ทธ๋ž˜ํ”„๋Š” ์•ˆ์ •์ ์œผ๋กœ ํ•™์Šต๋˜๋‚˜
173
+ MoFlow ์˜ one-step flow matching ์—์„œ๋Š” ์ด๋“์ด ์ œํ•œ์ "์ด๋ผ๋Š” ๊ฒƒ ์ž์ฒด๊ฐ€, ๋ฒ„๊ทธํŒ์˜
174
+ 0.695 ๊ฐ€ **์˜๋„ํ•œ ๋ฉ”์ปค๋‹ˆ์ฆ˜์ด ์•„๋‹ˆ๋ผ ์šฐ์—ฐํ•œ ์ •๊ทœํ™” ํšจ๊ณผ**์˜€์„ ์ˆ˜ ์žˆ์Œ์„ ์‹œ์‚ฌํ•˜๋Š”
175
+ ๋ฐœ๊ฒฌ์ด๋‹ค.
176
+
177
+ **โ‘ฃ MID / LED ๋Š” ๋‹ค๋ฅธ ์ฒ˜๋ฐฉ์„ ์“ด๋‹ค.** RES_CAP ์€ MoFlow ์ „์šฉ์ด๋‹ค. MID ๋Š”
178
+ `graph_gate_init` 0.01โ†’0.001 + warmup 100โ†’200 ์œผ๋กœ ๋ถ•๊ดด๋ฅผ ๋ง‰์•˜๊ณ , LED ๋Š” ๋ณ„๋„
179
+ ์ฒ˜๋ฐฉ ์—†์ด ํšŒ๋ณต ์ค‘์ด๋‹ค(LED sparse no-ฯƒ ๋Š” ์ด๋ฏธ ๊ธฐ์กด๊ฐ’ 0.886 โ†’ **0.800** ์œผ๋กœ ์•ž์„ฌ).
180
+
181
+ ---
182
+
183
+ ## 8. ์žฌํ˜„
184
+
185
+ ```bash
186
+ # MoFlow-NBA full SRA (ํ˜„์žฌ ์‹คํ–‰ ์„ค์ •)
187
+ SRA_EDGE_FIX=1 SRA_SOFT_START=1 SRA_RES_CAP=3.0 \
188
+ CUDA_VISIBLE_DEVICES=0 python fm_nba_graph_v6.py \
189
+ --cfg cfg/nba/cor_fm.yml --exp edgefix_cap3 \
190
+ --batch_size 192 --epochs 150 --fm_in_scaling --tied_noise \
191
+ --top_n_neighbors 5 --uncertainty_weight 0.01 --data_dir ./data/nba
192
+
193
+ # no-ฯƒ ablation ์€ ์ „์šฉ ์Šคํฌ๋ฆฝํŠธ ์‚ฌ์šฉ (use_sigma_gating=False ํ•˜๋“œ์ฝ”๋”ฉ)
194
+ SRA_EDGE_FIX=1 SRA_SOFT_START=1 SRA_RES_CAP=3.0 \
195
+ python fm_nba_graph_v6_nosigma.py --cfg cfg/nba/cor_fm.yml --exp ef_sparse_nosig \
196
+ --batch_size 192 --epochs 150 --fm_in_scaling --tied_noise --top_n_neighbors 5
197
+ ```
198
+
199
+ | ํ† ๊ธ€ | ๊ธฐ๋ณธ | ์—ญํ•  |
200
+ |---|---|---|
201
+ | `SRA_EDGE_FIX` | off | edge-index ์”ฌ ํ˜ผํ•ฉ ๋ฒ„๊ทธ ์ˆ˜์ • (**ํ•ต์‹ฌ**) |
202
+ | `SRA_RES_CAP` | 0 (off) | ๋…ธ๋“œ๋ณ„ residual norm ์ƒํ•œ โ€” MoFlow ์ƒ˜ํ”Œ๋ง ์•ˆ์ •ํ™” (**ํ•ต์‹ฌ**) |
203
+ | `SRA_SOFT_START` | off | out_proj zero-init + gate bias โˆ’4 (๋ณด์กฐ, ๋‹จ๋…์œผ๋กœ๋Š” ์‹คํŒจ) |
204
+ | `SRA_GATE_SCALE` | 1.0 | ๊ฒŒ์ดํŠธ ์ „์—ญ ์ถ•์†Œ (์‹คํŒจํ•œ ์ฒ˜๋ฐฉ, ์ฝ”๋“œ๋งŒ ์ž”์กด) |
docs/RELCAP.md ADDED
@@ -0,0 +1,181 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # `SRA_RES_CAP_REL` โ€” ์ƒ๋Œ€ residual cap
2
+
3
+ **ํ•œ ์ค„ ์š”์•ฝ**
4
+ edge ๋ฒ„๊ทธ๋ฅผ ๊ณ ์น˜์ž MoFlow ์ƒ˜ํ”Œ๋ง์ด ๋ฐœ์‚ฐํ–ˆ๋‹ค. ์ ˆ๋Œ€ cap(`SRA_RES_CAP`)์€ ๊ฐ’์„ ๋‚ฎ์ถฐ๋„
5
+ ๋ฐœ์‚ฐ ์‹œ์ ์„ ๋ฏธ๋ฃฐ ๋ฟ์ด์—ˆ๋‹ค. ์ƒํ•œ์„ **ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ norm ์— ๋น„๋ก€**ํ•˜๋„๋ก ๋ฐ”๊พธ์ž
6
+ (`โ€–resโ€– โ‰ค ratioยทโ€–origโ€–`) 13ํšŒ ํ‰๊ฐ€๊นŒ์ง€ ์•ˆ์ •์ ์œผ๋กœ ํ•˜๊ฐ•ํ•˜๋ฉฐ baseline ์„ ์•ž์„ฐ๋‹ค.
7
+
8
+ ---
9
+
10
+ ## 1. ์™œ ์ ˆ๋Œ€ cap ์œผ๋กœ๋Š” ๋ถ€์กฑํ•œ๊ฐ€
11
+
12
+ ๋ฐœ์‚ฐ์˜ ์›์ธ์€ MoFlow ์˜ **train/sample mismatch** ๋‹ค. ํ•™์Šต์€ ๋žœ๋ค timestep ํ•˜๋‚˜์—์„œ
13
+ ๊ทธ๋ž˜ํ”„๋ฅผ ํ•œ ๋ฒˆ๋งŒ ์ ์šฉํ•˜์ง€๋งŒ, ์ƒ˜ํ”Œ๋ง์€ 10 ์Šคํ… flow ์ ๋ถ„์—์„œ ๋งค ์Šคํ… ์ ์šฉํ•˜๊ณ  ๊ทธ
14
+ ์ถœ๋ ฅ์ด ๋‹ค์Œ ์Šคํ… ์ž…๋ ฅ์œผ๋กœ ๋˜๋จน์ž„๋œ๋‹ค โ†’ ์„ญ๋™์ด ๋ˆ„์ ๋œ๋‹ค.
15
+
16
+ ์ ˆ๋Œ€ cap ์€ `โ€–resโ€– โ‰ค c` ๋กœ **๊ณ ์ • ํฌ๊ธฐ**๋ฅผ ๊ฐ•์ œํ•œ๋‹ค. ๋ฌธ์ œ๋Š” ๋‘ ๊ฐ€์ง€๋‹ค.
17
+
18
+ 1. **ํ˜ธ์ŠคํŠธ ์Šค์ผ€์ผ ์˜์กด.** residual norm ์€ ๊ทธ ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ์ด ์“ฐ๋Š” ์Šค์ผ€์ผ ์œ„์—
19
+ ์žˆ๋‹ค. ์ธก์ •๊ฐ’: MoFlow ์—์„œ ๊ทธ๋ž˜ํ”„ residual โ‰ˆ **6.3**, ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ โ‰ˆ **16**.
20
+ MoFlow ์—์„œ ํŠœ๋‹ํ•œ `c = 3.0` ์ด ์ž„๋ฒ ๋”ฉ ์Šค์ผ€์ผ์ด ๋‹ค๋ฅธ MID/LED ์—์„œ๋Š” ์‚ฌ์‹ค์ƒ
21
+ ๋ฌด์—ฐ์‚ฐ์ด๊ฑฐ๋‚˜ ๋ฐ˜๋Œ€๋กœ ๊ณผ๋„ํ•œ ์ œ์•ฝ์ด ๋œ๋‹ค.
22
+ 2. **๋ˆ„์ ์„ ๊ฒฐ์ •ํ•˜๋Š” ๊ฒƒ์€ ์ ˆ๋Œ€ ํฌ๊ธฐ๊ฐ€ ์•„๋‹ˆ๋ผ ๋น„์œจ.** ์Šคํ…๋‹น ์ƒํƒœ๊ฐ€
23
+ `x โ† x + res` ๋กœ ๊ฐฑ์‹ ๋  ๋•Œ ๋ฐœ์‚ฐ ์—ฌ๋ถ€๋ฅผ ์ง€๋ฐฐํ•˜๋Š” ๊ฒƒ์€ `โ€–resโ€–/โ€–xโ€–` ๋‹ค. ์ ˆ๋Œ€ cap ์€
24
+ ์ด ๋น„์œจ์„ ์ง์ ‘ ํ†ต์ œํ•˜์ง€ ๋ชปํ•œ๋‹ค.
25
+
26
+ ์‹ค์ œ๋กœ ์ ˆ๋Œ€ cap ์€ ๊ฐ’์„ ๋‚ฎ์ถœ์ˆ˜๋ก **๋ฐœ์‚ฐ์ด ๋ฏธ๋ค„์งˆ ๋ฟ ์‚ฌ๋ผ์ง€์ง€ ์•Š์•˜๋‹ค**:
27
+
28
+ | ์ ˆ๋Œ€ cap | ๋ฐœ์‚ฐ ์‹œ์  |
29
+ |---|---|
30
+ | 6.0 | eval 4 |
31
+ | 3.0 | eval 5 |
32
+ | 1.0 | eval 9 (์ดํ›„ 1.15~1.61 ์ง„๋™) |
33
+ | 0.5 | 12ํšŒ์ฐจ๋ถ€ํ„ฐ ๋ฐ˜๋“ฑ (0.873 โ†’ 0.922 โ†’ 0.958) |
34
+
35
+ ---
36
+
37
+ ## 2. ๊ตฌํ˜„
38
+
39
+ ```python
40
+ # SRA_RES_CAP_REL: ๋…ธ๋“œ๋ณ„ ์ƒํ•œ์„ ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ norm ์— ๋น„๋ก€ํ•ด ์ •ํ•œ๋‹ค.
41
+ _rel = float(os.environ.get('SRA_RES_CAP_REL', 0.0) or 0.0)
42
+ if _rel > 0:
43
+ lim = _rel * orig.norm(dim=-1, keepdim=True) # [N, 1] ๋…ธ๋“œ๋ณ„ ์ƒํ•œ
44
+ rn = res.norm(dim=-1, keepdim=True)
45
+ res = res * torch.where(rn > lim, lim / rn.clamp_min(1e-6),
46
+ torch.ones_like(rn))
47
+ out = orig + res
48
+ ```
49
+
50
+ - `orig` = ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ, `res` = ๊ทธ๋ž˜ํ”„ ์„ญ๋™. ์ƒํ•œ์ด **๋…ธ๋“œ๋งˆ๋‹ค** ์ž๊ธฐ ์ž„๋ฒ ๋”ฉ ํฌ๊ธฐ์—
51
+ ๋น„๋ก€ํ•ด ์ •ํ•ด์ง€๋ฏ€๋กœ ์Šค์ผ€์ผ ๋ฌด๊ด€ํ•˜๋‹ค.
52
+ - `torch.where` ๋ฅผ ์“ฐ๋Š” ์ด์œ ๋Š” ยง4 ์ฐธ์กฐ (์ด์ „ ๊ตฌํ˜„์˜ ์น˜๋ช…์  ๋ฒ„๊ทธ).
53
+ - ๊ธฐ๋ณธ๊ฐ’ `0` = ๊บผ์ง. ์ผœ์ง€ ์•Š์œผ๋ฉด ์›๋ณธ๊ณผ ๋™์ผํ•˜๊ฒŒ ๋™์ž‘ํ•œ๋‹ค.
54
+
55
+ ๊ฒ€์ฆ (A=11, K=4):
56
+
57
+ | ์„ค์ • | out_proj grad | โ€–resโ€–max | โ€–resโ€–/โ€–origโ€– max |
58
+ |---|---|---|---|
59
+ | ๋ฌด์ œํ•œ | 1,716,997 | 9.009 | **0.557** |
60
+ | ์ ˆ๋Œ€ cap 1.0 | 255,169 | 1.000 | 0.071 |
61
+ | relcap 0.10 | 408,742 | 1.765 | **0.100** โœ“ |
62
+ | relcap 0.03 | 122,623 | 0.529 | **0.030** โœ“ |
63
+
64
+ ๋น„์œจ์ด ์ง€์ •๊ฐ’์œผ๋กœ ์ •ํ™•ํžˆ ์ œํ•œ๋˜๊ณ  gradient ๋„ ์ •์ƒ์ด๋‹ค.
65
+
66
+ ---
67
+
68
+ ## 3. ๊ฒฐ๊ณผ (MoFlow-NBA, full SRA, `SRA_EDGE_FIX=1`)
69
+
70
+ eval ๋ณ„ min-ADEโ‚‚โ‚€ @4.0s:
71
+
72
+ | eval | 1 | 2 | 3 | 4 | 5 | 6 | 7 | 8 | 9 | 10 | 11 | 12 | 13 |
73
+ |---|---|---|---|---|---|---|---|---|---|---|---|---|---|
74
+ | **relcap 0.03** | 1.190 | 1.010 | 0.979 | 1.019 | 0.971 | 0.941 | 0.912 | 0.894 | 0.894 | 0.857 | 0.863 | 0.863 | **0.840** |
75
+ | **relcap 0.10** | 1.128 | 1.017 | 0.988 | 0.954 | 0.950 | 0.944 | 0.958 | 0.936 | 0.880 | 0.889 | 0.857 | 0.862 | 0.876 |
76
+ | cap 0.5 (์ ˆ๋Œ€) | 1.155 | 1.032 | 1.018 | 0.990 | 0.978 | 0.907 | 0.937 | 0.901 | 0.931 | 0.873 | 0.875 | 0.922 | 0.958 |
77
+ | cap 1.0 (์ ˆ๋Œ€) | 1.146 | 1.013 | 0.985 | 1.020 | 0.978 | 1.006 | 1.170 | 0.966 | **1.309** | 1.378 | 1.150 | 1.344 | 1.370 |
78
+ | *baseline(์ฐธ๊ณ )* | 1.163 | 1.021 | 0.961 | 0.964 | 0.952 | 0.955 | 0.953 | 0.895 | โ€” | โ€” | โ€” | โ€” | โ€” |
79
+
80
+ ํ˜„์žฌ best (ADE/FDE ๋Š” **๊ฐ™์€ ํ‰๊ฐ€ ์‹œ์ **์—์„œ ์ง์ง€์Œ):
81
+
82
+ | ์„ค์ • | best ADE / FDE | @eval | ์ง„ํ–‰ |
83
+ |---|---|---|---|
84
+ | **relcap 0.03** | **0.8399 / 1.0025** | 13 | 2368/25500 (9 %) |
85
+ | relcap 0.10 | 0.8568 / 1.0441 | 11 | 2349/25500 |
86
+ | cap 0.5 | 0.8732 / 1.0531 | 10 | 2379/25500 |
87
+ | cap 1.0 | 0.9662 / 1.2508 | 8 | 3390/25500 โ€” ๋ฐœ์‚ฐ |
88
+
89
+ ๊ด€์ฐฐ:
90
+
91
+ - **์›๋ž˜ ๋ฐœ์‚ฐ ์ง€์ (eval 4~5)์„ ์ƒ๋Œ€ cap 2์ข… ๋ชจ๋‘ ํ†ต๊ณผ**ํ–ˆ๋‹ค. ์ ˆ๋Œ€ cap 1.0 ๋„ 8ํšŒ์ฐจ๊นŒ์ง€๋Š”
92
+ ๋ฉ€์ฉกํ•ด ๋ณด์˜€์œผ๋ฏ€๋กœ 4~5ํšŒ ํ†ต๊ณผ๋งŒ์œผ๋กœ๋Š” ๋ถ€์กฑํ•œ๋ฐ, 13ํšŒ๊นŒ์ง€ ์œ ์ง€๋œ ๊ฒƒ์€ ๋‹ค๋ฅธ ์ˆ˜์ค€์˜ ๊ทผ๊ฑฐ๋‹ค.
93
+ - **baseline(0.895) ์„ ์•ž์„ฐ๋‹ค.** relcap 0.03 ์ด 0.840 ์œผ๋กœ โˆ’0.055.
94
+ - ์ ˆ๋Œ€ cap 0.5 ๋Š” 10ํšŒ์ฐจ 0.873 ์ดํ›„ **0.922 โ†’ 0.958 ๋กœ 3ํšŒ ์—ฐ์† ์ƒ์Šน** โ€” ์ƒ๋Œ€ cap ๊ณผ
95
+ ๋‹ฌ๋ฆฌ ๋ถˆ์•ˆ์ • ์กฐ์ง์ด ๋ณด์ธ๋‹ค.
96
+ - ๋น„์œจ์„ 0.10 โ†’ 0.03 ์œผ๋กœ ๋” ์กฐ์—ฌ๋„ ์„ฑ๋Šฅ์ด ๋‚˜๋น ์ง€์ง€ ์•Š์•˜๋‹ค. ์ฆ‰ **๋ฌด์ œํ•œ ์ƒํƒœ์˜ ๋น„์œจ
97
+ 0.557 ์€ ํ•„์š” ์ด์ƒ์œผ๋กœ ํฌ๊ณ , ๊ทธ ๊ผฌ๋ฆฌ๊ฐ€ ๋ฐœ์‚ฐ์„ ์œ ๋ฐœ**ํ–ˆ๋‹ค๋Š” ํ•ด์„๊ณผ ์ผ์น˜ํ•œ๋‹ค.
98
+
99
+ ---
100
+
101
+ ## 4. โš ๏ธ ๊ฐ™์ด ๊ณ ์นœ ๊ฒƒ โ€” ์ด์ „ cap ๊ตฌํ˜„์˜ ์น˜๋ช…์  ๋ฒ„๊ทธ
102
+
103
+ ์ฒ˜์Œ ์ž‘์„ฑํ•œ cap ์€ ์ด๋žฌ๋‹ค:
104
+
105
+ ```python
106
+ res = res * (rn.clamp(max=cap) / (rn + 1e-6)) # ์ž˜๋ชป๋จ
107
+ ```
108
+
109
+ ๋‘ ๊ฐ€์ง€๊ฐ€ ํ‹€๋ ธ๋‹ค.
110
+
111
+ **(1) `res = 0` ์—์„œ gradient ๊ฐ€ ์ •ํ™•ํžˆ 0 ์ด ๋œ๋‹ค.** ์Šค์ผ€์ผ์ด `0 / 1e-6 = 0` ์ด ๋˜๊ณ 
112
+ Jacobian ๋„ 0 ์ด๋‹ค. `SRA_SOFT_START`(out_proj zero-init) ์™€ ํ•จ๊ป˜ ์ผœ๋ฉด `out_proj` ๊ฐ€
113
+ 0 ์— ์˜๊ตฌํžˆ ๊ฐ‡ํ˜€ **๊ทธ๋ž˜ํ”„๊ฐ€ ์ „ํ˜€ ํ•™์Šต๋˜์ง€ ์•Š๋Š”๋‹ค**. NaN ์ด ์•„๋‹ˆ๋ผ ์กฐ์šฉํžˆ ์ฃฝ๋Š”๋‹ค.
114
+
115
+ ์‹ค์ธก (out_proj gradient ํ•ฉ):
116
+
117
+ | ์„ค์ • | grad |
118
+ |---|---|
119
+ | SOFT_START ๋งŒ | 78,751 |
120
+ | RES_CAP ๋งŒ | 765,506 |
121
+ | **SOFT_START + RES_CAP** | **0.00** โ† ํ•™์Šต ๋ถˆ๊ฐ€ |
122
+
123
+ ์ด ์กฐํ•ฉ์œผ๋กœ ๋Œ๋ฆฐ ์‹คํ–‰๋“ค์˜ ์ฒดํฌํฌ์ธํŠธ๋ฅผ ์—ด์–ด๋ณด๋‹ˆ **58 epoch ๋’ค์—๋„
124
+ `future_graph.out_proj` ๊ฐ€ ์ •ํ™•ํžˆ `0.000e+00`** ์ด์—ˆ๋‹ค. ์ฆ‰ ๊ทธ ์‹คํ–‰๋“ค์€ host ๋‹จ๋…
125
+ (=baseline) ์„ ์ธก์ •ํ•œ ๊ฒƒ์ด๊ณ , "cap ์ด ๋ฐœ์‚ฐ์„ ํ•ด๊ฒฐํ–ˆ๋‹ค"๋˜ ๊ฒฐ๋ก ์€ **๊ทธ๋ž˜ํ”„๊ฐ€ ๊บผ์ ธ ์žˆ์–ด
126
+ ๋ฐœ์‚ฐํ•  ๊ฒƒ์ด ์—†์—ˆ์„ ๋ฟ**์ด์—ˆ๋‹ค. ํ•ด๋‹น ๊ฒฐ๊ณผ 4๊ฑด์€ ํ๊ธฐํ–ˆ๋‹ค.
127
+
128
+ **(2) cap ๋ฏธ๋งŒ ๊ฐ’๋„ ๋ถ€๋‹นํ•˜๊ฒŒ ์ถ•์†Œ๋œ๋‹ค.** `rn/(rn+1e-6)` ์€ `rn` ์ด ์ž‘์„์ˆ˜๋ก 1 ์—์„œ
129
+ ๋ฉ€์–ด์ง„๋‹ค โ€” `โ€–resโ€–=1e-5` ์—์„œ **0.909 ๋ฐฐ**.
130
+
131
+ `torch.where` ๋กœ ๋ฐ”๊พธ๋ฉด ๋‘˜ ๋‹ค ํ•ด๊ฒฐ๋œ๋‹ค:
132
+
133
+ | ๊ฒ€์ฆ | old | new |
134
+ |---|---|---|
135
+ | `res=0` ์—์„œ gradient | 0.0000 | **31.30** |
136
+ | โ€–resโ€–=1e-5 ์ผ ๋•Œ ์Šค์ผ€์ผ | 0.909 | **1.000** |
137
+ | cap ์ดˆ๊ณผ ์‹œ ์ƒํ•œ | 3.0 | 3.0 (๋™์ผ) |
138
+ | ์ •์ƒ ๊ตฌ๊ฐ„ gradient ์ฐจ์ด | โ€” | ์ƒ๋Œ€์ฐจ 6e-7 |
139
+
140
+ ---
141
+
142
+ ## 5. ์‚ฌ์šฉ๋ฒ•
143
+
144
+ ```bash
145
+ # ๊ถŒ์žฅ ์„ค์ • (MoFlow-NBA)
146
+ SRA_EDGE_FIX=1 SRA_RES_CAP_REL=0.03 \
147
+ CUDA_VISIBLE_DEVICES=0 python fm_nba_graph_v6.py \
148
+ --cfg cfg/nba/cor_fm.yml --exp v3_rel003 \
149
+ --batch_size 192 --epochs 150 --fm_in_scaling --tied_noise \
150
+ --top_n_neighbors 5 --uncertainty_weight 0.01 --data_dir ./data/nba
151
+ ```
152
+
153
+ | ํ† ๊ธ€ | ๊ธฐ๋ณธ | ์—ญํ•  |
154
+ |---|---|---|
155
+ | `SRA_EDGE_FIX` | off | edge-index ์”ฌ ํ˜ผํ•ฉ ๋ฒ„๊ทธ ์ˆ˜์ • (**ํ•„์ˆ˜**) |
156
+ | `SRA_RES_CAP_REL` | 0 (off) | **๋น„์œจ ์ƒํ•œ** `โ€–resโ€– โ‰ค rยทโ€–origโ€–` โ€” ๊ถŒ์žฅ |
157
+ | `SRA_RES_CAP` | 0 (off) | ์ ˆ๋Œ€ ์ƒํ•œ โ€” MoFlow ์—์„œ ์—ด๋“ฑํ•จ์ด ํ™•์ธ๋จ |
158
+ | `SRA_SOFT_START` | off | **RES_CAP ๊ณผ ํ•จ๊ป˜ ์“ฐ์ง€ ๋ง ๊ฒƒ** (ยง4). ๋‹จ๋…์œผ๋กœ๋„ ๋ฐœ์‚ฐ ๋ชป ๋ง‰์Œ |
159
+ | `SRA_GATE_SCALE` | 1.0 | ์ „์—ญ ์ถ•์†Œ โ€” ๊ฐ€์žฅ ๋นจ๋ฆฌ ๋ฐœ์‚ฐ(3ํšŒ์ฐจ), ํ๊ธฐ |
160
+
161
+ ---
162
+
163
+ ## 6. ๋ฏธํ™•์ • / ํ•œ๊ณ„
164
+
165
+ **โ‘  ์ตœ์ข… ์ˆ˜๋ ด๊ฐ’์€ ์•„์ง ๋ชจ๋ฅธ๋‹ค.** ํ˜„์žฌ 9 %(2368/25500) ์—์„œ 0.840 ์ด๋‹ค.
166
+ ๊ธฐ์กด(๋ฒ„๊ทธํŒ) SRA = **0.695**, baseline = **0.703**. ๋‚จ์€ 91 % ์™€ cosine LR ๊ฐ์‡ ์—์„œ
167
+ ๋” ๋‚ด๋ ค๊ฐ€์•ผ ํ•˜๋ฉฐ, **0.695 ์— ๋„๋‹ฌํ•˜์ง€ ๋ชปํ•  ๊ฐ€๋Šฅ์„ฑ์€ ์—ฌ์ „ํžˆ ์—ด๋ ค ์žˆ๋‹ค.**
168
+
169
+ **โ‘ก ๋น„์œจ๊ฐ’ ํŠœ๋‹์ด ๋๋‚˜์ง€ ์•Š์•˜๋‹ค.** 0.03 ๊ณผ 0.10 ์ด ๋น„์Šทํ•˜๊ณ  0.03 ์ด ๊ทผ์†Œ ์šฐ์œ„๋‹ค.
170
+ ๋” ์กฐ์ธ ๊ฐ’(0.01)์ด ๋‚˜์„์ง€, ์•„๋‹ˆ๋ฉด 0.03 ์ด ์ด๋ฏธ ๊ณผ๋„ํ•œ ์ œ์•ฝ์ธ์ง€๋Š” ๋ฏธ๊ฒ€์ฆ์ด๋‹ค.
171
+ cap ์„ ์กฐ์ผ์ˆ˜๋ก ๋ฐœ์‚ฐ์€ ๋ง‰ํžˆ์ง€๋งŒ ๊ทธ๋ž˜ํ”„ ๊ธฐ์—ฌ๋„ ์ค„์–ด๋“œ๋ฏ€๋กœ, **"๋ฐœ์‚ฐ ์•ˆ ํ•˜๋ฉด์„œ
172
+ baseline ๋ณด๋‹ค ๋‚˜์€" ๊ตฌ๊ฐ„์ด ์‹ค์ œ๋กœ ์กด์žฌํ•˜๋Š”์ง€**๊ฐ€ ์ตœ์ข… ์งˆ๋ฌธ์ด๋‹ค. ํ˜„์žฌ 0.840 vs
173
+ baseline 0.895 ๋Š” ๊ทธ ๊ตฌ๊ฐ„์ด ์กด์žฌํ•œ๋‹ค๋Š” ์ฒซ ์ฆ๊ฑฐ๋‹ค.
174
+
175
+ **โ‘ข MoFlow ์ „์šฉ ์ฒ˜๋ฐฉ์ด๋‹ค.** MIDยทLED ๋Š” iterative denoiser ๋ผ ์ด ๋ฐœ์‚ฐ์ด ์—†๊ณ ,
176
+ `RES_CAP*` ์„ ์“ฐ์ง€ ์•Š๋Š”๋‹ค (MID ๋Š” gate_init/warmup, LED ๋Š” ์ฒ˜๋ฐฉ ์—†์Œ).
177
+ ๋‹ค๋ฅธ ํ˜ธ์ŠคํŠธ์— ์ ์šฉํ•˜๋ ค๋ฉด ๊ทธ ํ˜ธ์ŠคํŠธ์˜ residual/embedding ๋น„์œจ์„ ๋จผ์ € ์ธก์ •ํ•ด์•ผ ํ•œ๋‹ค.
178
+
179
+ **โ‘ฃ ์ ˆ๋Œ€ cap ์ด ํ•ญ์ƒ ๋‚˜์˜๋‹ค๊ณ  ๋‹จ์ •ํ•  ์ˆ˜๋Š” ์—†๋‹ค.** cap 0.5 ๋Š” ์•„์ง ๋ฐœ์‚ฐํ•˜์ง€ ์•Š์•˜๊ณ 
180
+ ๋ฐ˜๋“ฑ ์กฐ์ง๋งŒ ๋ณด์ธ๋‹ค. ์ƒ๋Œ€ cap ์˜ ์šฐ์œ„๋Š” ํ˜„์žฌ 13ํšŒ ํ‰๊ฐ€ ๊ธฐ์ค€์˜ ๊ด€์ฐฐ์ด๋ฉฐ, ์™„์ฃผ๊นŒ์ง€
181
+ ๊ฐ€์•ผ ํ™•์ •๋œ๋‹ค.