| # MoFlow ์ต์ ๋ณธ โ ์ด์ ๋ฒ์ ๋๋น ๋ฌด์์ด ๋ฐ๋์๋ |
|
|
| ๋น๊ต ๊ธฐ์ค: HF `po03087/sra-trajectory-code` ์ปค๋ฐ `0d52562` โ `37c61d4` |
|
|
| ๋ฐ๋ MoFlow ํ์ผ์ **2๊ฐ, ์ด 88์ค**์ด๋ค. ๊ทธ ์ค ์ค์ ๋ก์ง์ 5์ค์ด๊ณ ๋๋จธ์ง๋ ์ฃผ์์ด๋ค. |
|
|
| | ํ์ผ | ๋ณ๊ฒฝ | ์ฑ๊ฒฉ | |
| |---|---|---| |
| | `models/graph_interaction_nba.py` | +23 | **๋ฒ๊ทธ ์์ ** (edge index) | |
| | `models/graph_interaction_nba_v6.py` | +65 | **์ ์์ ํ ์ต์
** 3์ข
| |
|
|
| > **๋ชจ๋ ์ ๋์์ ํ๊ฒฝ๋ณ์ ํ ๊ธ ๋ค์ ์๊ณ ๊ธฐ๋ณธ๊ฐ์ ์ ๋ถ off๋ค.** |
| > ํ ๊ธ์ ์ ์ผ๋ฉด ์ด์ ์ฝ๋์ ๋นํธ ๋จ์๋ก ๋์ผํ๊ฒ ๋์ํ๋ค โ ๊ธฐ์กด ์ฒดํฌํฌ์ธํธ์ |
| > ์งํ ์ค์ธ ์คํ์ด ๊ทธ๋๋ก ์ฌํ๋๋ค. |
|
|
| --- |
|
|
| ## 1. `graph_interaction_nba.py` โ edge index ์ฌ ํผํฉ ๋ฒ๊ทธ |
|
|
| ### ๋ฐ๋ ์ฝ๋ |
|
|
| ```diff |
| def _make_batched_edge_index(self, num_scenes): |
| single = self._single_edge_index # [2, E0] |
| offsets = torch.arange(num_scenes, ...) * self.num_agents |
| batched = single.unsqueeze(0).expand(num_scenes,-1,-1) + offsets.view(-1,1,1) |
| - return batched.reshape(2, -1) # [2, S*E0] |
| + if os.environ.get('SRA_EDGE_FIX', '') not in ('', '0', 'false', 'False'): |
| + # [S,2,E0] -> [2,S,E0] -> [2,S*E0]: src/dst ์ถ์ ๋จผ์ ์์ผ๋ก ์ฎ๊ธด๋ค |
| + return batched.permute(1, 0, 2).reshape(2, -1) |
| + return batched.reshape(2, -1) # legacy (์ฌ ํผํฉ) |
| ``` |
|
|
| ### ์ ํ๋ ธ๋ |
|
|
| `batched` ์ shape ๋ `[S, 2, E0]` ์ด๊ณ ๋ฉ๋ชจ๋ฆฌ ๋ฐฐ์น๋ |
|
|
| ``` |
| s0_src s0_dst s1_src s1_dst s2_src s2_dst ... |
| ``` |
|
|
| ์ด๋ค. ์ฌ๊ธฐ์ `.reshape(2, -1)` ์ ํ๋ฉด ์ ์ ๋ฐ์ด row 0, ๋ค ์ ๋ฐ์ด row 1 ์ด ๋๋ฏ๋ก |
| **row 0 = "์์ชฝ ์ฌ๋ค์ src+dst ์ ๋ถ", row 1 = "๋ค์ชฝ ์ฌ๋ค์ src+dst ์ ๋ถ"** ๊ฐ ๋๋ค. |
| src/dst ํ์ด ์๋๋ผ **์ฌ ๋ธ๋ก**์ผ๋ก ์๋ฆฐ ๊ฒ์ด๋ค. |
|
|
| ### ์ธก์ ๋ ์ํฅ (A=11) |
|
|
| | ์งํ | ์์ ์ | ์์ ํ | |
| |---|---|---| |
| | ๊ฐ์ ์ฌ ์์ ๋จธ๋ฌด๋ edge ๋น์จ | **0 %** | **100 %** | |
| | ๋ค์ด์ค๋ edge ๊ฐ ํ๋๋ ์๋ ๋
ธ๋ | **์ ๋ฐ** | 0 | |
| | ๋๋จธ์ง ๋
ธ๋์ degree | ์๋์ **2๋ฐฐ** | ์ ์ | |
| | top-N ์ด์์ด ์ฌ๋ฐ๋ฅด๊ฒ ๋ฌถ์ด๋ ๋น์จ | 33 % | 100 % | |
|
|
| ์ฆ "์์ธก๋ ๋ฏธ๋ ์์์ ์ฌ ๋ด๋ถ top-N ์ด์์ ๊ณ ๋ฅธ๋ค"๋ SRA ์ ์ค๋ช
๊ณผ ์ค์ ์ฝ๋๊ฐ |
| ๋ฌ๋๋ค. ๋
ผ๋ฌธ ์์น(MID 0.957 / LED 0.778 / MoFlow 0.695)๋ ์ ๋ถ ์ด ๋ฒ๊ทธ ์ํ์์ |
| ๋์จ ๊ฐ์ด๋ค. |
|
|
| --- |
|
|
| ## 2. `graph_interaction_nba_v6.py` โ ์์ ํ ์ต์
3์ข
|
| |
| edge ๋ฅผ ๊ณ ์น์ **MoFlow ๋ง** ์ ๋ฌธ์ ๊ฐ ์๊ฒผ๋ค. ํ์ต loss ๋ ๋จ์กฐ ๊ฐ์ํ๋๋ฐ |
| **์ํ๋ง์ด ๋ฐ์ฐ**ํ๋ค. ์์ธ์ train/sample mismatch ๋ค. |
| |
| - ํ์ต: ๋๋ค timestep ํ๋์์ ๊ทธ๋ํ๋ฅผ **ํ ๋ฒ** ์ ์ฉ |
| - ์ํ๋ง: 10 ์คํ
flow ์ ๋ถ์์ **๋งค ์คํ
** ์ ์ฉํ๊ณ ๊ทธ ์ถ๋ ฅ์ด ๋ค์ ์คํ
์
๋ ฅ์ผ๋ก ๋๋จน์ |
| |
| ๋ฒ๊ทธ ์ํ์์๋ ๋
ธ๋ ์ ๋ฐ์ด ๊ณ ์๋ผ ์ญ๋์ด ์ฝํด์ ์ด ๋์ ์ด ๋๋ฌ๋์ง ์์๋ค. |
| edge ๋ฅผ ๊ณ ์ณ ๋ชจ๋ ๋
ธ๋๊ฐ ์ด์์ ๋ฐ์ ์ญ๋์ด ์ปค์ง๊ณ 10 ์คํ
์ ๊ฑธ์ณ ๊ธฐํ๊ธ์๋ก ์์ธ๋ค. |
| |
| MIDยทLED ๋ iterative denoiser ๋ผ ์ด ํ์์ด ์๋ค. **์๋ ์ต์
์ MoFlow ์ ์ฉ ์ฒ๋ฐฉ์ด๋ค.** |
| |
| ### 2.1 `SRA_SOFT_START` โ ์ด๊ธฐ ์ถฉ๊ฒฉ ์ ๊ฑฐ |
| |
| ```python |
| if os.environ.get('SRA_SOFT_START', ...): |
| nn.init.zeros_(self.out_proj.weight); nn.init.zeros_(self.out_proj.bias) |
| nn.init.constant_(gate_proj_linear.bias, SRA_GATE_BIAS) # ๊ธฐ๋ณธ -4, sigmoid(-4)โ0.018 |
| ``` |
| |
| MID ๊ฐ ์๋ ์ฐ๋ ๋ฐฉ์(zero-init + warmup gate)์ V6 ๋ก ์ฎ๊ธด ๊ฒ. ๋ ๋ค **ํ์ต ๊ฐ๋ฅ** |
| ํ๊ฒ ๋จ์ ์์ด์ ๊ณ ์ ์ถ์์ ๋ฌ๋ฆฌ ๊ทธ๋ํ๊ฐ ๋์ค์ ์ ๊ฐ๋๊น์ง ์๋ ์ ์๋ค. |
| **๋จ๋
์ผ๋ก๋ ๋ฐ์ฐ์ ๋ง์ง ๋ชปํ๋ค.** |
| |
| ### 2.2 `SRA_GATE_SCALE` โ ์ ์ญ ์ถ์ |
| |
| ```python |
| res = SRA_GATE_SCALE * gate * self.out_proj(nodes) # ๊ธฐ๋ณธ 1.0 = off |
| ``` |
| |
| ๊ฐ์ฅ ๋จ์ํ ์ฒ๋ฐฉ์ด์ง๋ง **๊ฐ์ฅ ๋นจ๋ฆฌ ๋ฐ์ฐํ๋ค(3ํ์ฐจ). ํ๊ธฐ.** |
| |
| ### 2.3 `SRA_RES_CAP` / `SRA_RES_CAP_REL` โ residual norm ์ํ |
|
|
| ```python |
| _cap = float(os.environ.get('SRA_RES_CAP', 0.0) or 0.0) # ์ ๋ ์ํ |
| if _cap > 0: |
| rn = res.norm(dim=-1, keepdim=True) |
| res = res * torch.where(rn > _cap, _cap / rn.clamp_min(1e-6), torch.ones_like(rn)) |
| |
| _rel = float(os.environ.get('SRA_RES_CAP_REL', 0.0) or 0.0) # ์๋ ์ํ โ
|
| if _rel > 0: |
| lim = _rel * orig.norm(dim=-1, keepdim=True) # ๋
ธ๋๋ง๋ค ์๊ธฐ ์๋ฒ ๋ฉ ํฌ๊ธฐ์ ๋น๋ก |
| rn = res.norm(dim=-1, keepdim=True) |
| res = res * torch.where(rn > lim, lim / rn.clamp_min(1e-6), torch.ones_like(rn)) |
| out = orig + res |
| ``` |
|
|
| **์ ๋ cap ์ ๊ฐ์ ๋ฎ์ถฐ๋ ๋ฐ์ฐ ์์ ์ ๋ฏธ๋ฃฐ ๋ฟ์ด์๋ค:** |
|
|
| | ์ ๋ cap | ๊ฒฐ๊ณผ | |
| |---|---| |
| | 6.0 | 4ํ์ฐจ ๋ฐ์ฐ | |
| | 3.0 | 5ํ์ฐจ ๋ฐ์ฐ | |
| | 1.0 | 8ํ์ฐจ best 0.9662 โ ์ดํ ๋ฐ์ฐ, ์ต๊ทผ 1.41 | |
| | 0.5 | 10ํ์ฐจ best 0.8732 โ 12ํ์ฐจ๋ถํฐ ๋ฐ๋ฑ, ์ต๊ทผ 1.41 | |
|
|
| ์ด์ ๋ ๋ ๊ฐ์ง๋ค. (a) residual norm ์ ํธ์คํธ ์๋ฒ ๋ฉ ์ค์ผ์ผ ์์ ์์ด์ MoFlow ์์ |
| ํ๋ํ ๊ฐ์ด MID/LED ์์๋ ์๋ฏธ๊ฐ ๋ฌ๋ผ์ง๋ค. (b) `x โ x + res` ์ ๋์ ์ ์ง๋ฐฐํ๋ ๊ฒ์ |
| ์ ๋ ํฌ๊ธฐ๊ฐ ์๋๋ผ **๋น์จ `โresโ/โxโ`** ์ธ๋ฐ ์ ๋ cap ์ ์ด๊ฑธ ์ง์ ํต์ ํ์ง ๋ชปํ๋ค. |
|
|
| ์๋ cap ์ ๋
ธ๋๋ง๋ค `โresโ โค rยทโorigโ` ๋ก **๋น์จ์ ์ง์ ** ๋ฌถ๋๋ค. ์ด๊ฒ๋ง ์ด์๋จ์๋ค. |
|
|
| ### โ ๏ธ ๊ฐ์ด ๊ณ ์น ๊ฒ โ ์ด์ cap ๊ตฌํ์ dead-gradient ๋ฒ๊ทธ |
|
|
| ์ฒ์ ์ด cap ์ `res * (rn.clamp(max=cap) / (rn + 1e-6))` ์ด์๋๋ฐ ๋ ๊ฐ์ง๊ฐ ํ๋ ธ๋ค. |
|
|
| 1. **`res = 0` ์์ gradient ๊ฐ ์ ํํ 0.** ์ค์ผ์ผ์ด `0/1e-6 = 0` ์ด๊ณ Jacobian ๋ 0. |
| `SRA_SOFT_START`(out_proj zero-init)์ ๊ฐ์ด ์ผ๋ฉด out_proj ๊ฐ 0 ์ ์๊ตฌํ ๊ฐํ |
| **๊ทธ๋ํ๊ฐ ์ ํ ํ์ต๋์ง ์๋๋ค.** NaN ์ด ์๋๋ผ ์กฐ์ฉํ ์ฃฝ๋๋ค. |
|
|
| | ์ค์ | out_proj gradient ํฉ | |
| |---|---| |
| | SOFT_START ๋ง | 78,751 | |
| | RES_CAP ๋ง | 765,506 | |
| | **SOFT_START + RES_CAP** | **0.00** โ ํ์ต ๋ถ๊ฐ | |
| |
| ์ค์ ๋ก ๊ทธ ์กฐํฉ์ผ๋ก ๋๋ฆฐ ์คํ์ ์ฒดํฌํฌ์ธํธ๋ 58 epoch ๋ค์๋ |
| `future_graph.out_proj` ๊ฐ ์ ํํ `0.000000e+00` ์ด์๋ค. ์ฆ ๊ทธ ์คํ๋ค์ host |
| ๋จ๋
(=baseline)์ ์ธก์ ํ ๊ฒ์ด๊ณ , "cap ์ด ๋ฐ์ฐ์ ํด๊ฒฐํ๋ค"๋ ๊ฒฐ๋ก ์ ๋ฌดํจ์๋ค. |
| ํด๋น ๊ฒฐ๊ณผ 4๊ฑด์ ํ๊ธฐํ๋ค. |
| |
| 2. **cap ๋ฏธ๋ง ๊ฐ๋ ๋ถ๋นํ๊ฒ ์ถ์๋๋ค.** `rn/(rn+1e-6)` ์ `rn` ์ด ์์์๋ก 1 ์์ |
| ๋ฉ์ด์ง๋ค โ `โresโ=1e-5` ์์ **0.909 ๋ฐฐ**. |
| |
| `torch.where` ๋ก ๋ฐ๊พธ๋ฉด ๋ ๋ค ํด๊ฒฐ๋๋ค. |
| |
| | ๊ฒ์ฆ | old | new | |
| |---|---|---| |
| | `res=0` ์์ gradient | 0.0000 | **31.30** | |
| | `โresโ=1e-5` ์ผ ๋ ์ค์ผ์ผ | 0.909 | **1.000** | |
| | cap ์ด๊ณผ ์ ์ํ | 3.0 | 3.0 (๋์ผ) | |
| | ์ ์ ๊ตฌ๊ฐ gradient ์ฐจ์ด | โ | ์๋์ฐจ 6e-7 | |
| |
| --- |
| |
| ## 3. ํ ๊ธ ์์ฝ |
| |
| | ํ ๊ธ | ๊ธฐ๋ณธ | ์ญํ | ํ์ | |
| |---|---|---|---| |
| | `SRA_EDGE_FIX` | off | edge index ์ฌ ํผํฉ ์์ | **ํ์** | |
| | `SRA_RES_CAP_REL` | 0 (off) | `โresโ โค rยทโorigโ` ๋น์จ ์ํ | **์ ์ผํ ์์กด ์ฒ๋ฐฉ** | |
| | `SRA_RES_CAP` | 0 (off) | ์ ๋ ์ํ | ๋ ๊ฐ ๋ชจ๋ ๋ฐ์ฐ, ์ด๋ฑ | |
| | `SRA_SOFT_START` | off | out_proj zero-init + gate bias | ๋จ๋
๋ถ์ถฉ๋ถ. **RES_CAP ๊ณผ ๋ณ์ฉ ๊ธ์ง** (ยง2.3) | |
| | `SRA_GATE_BIAS` | โ4.0 | SOFT_START ์ gate bias | โ | |
| | `SRA_GATE_SCALE` | 1.0 (off) | ์ ์ญ ์ถ์ | 3ํ์ฐจ ๋ฐ์ฐ, ํ๊ธฐ | |
|
|
| --- |
|
|
| ## 4. ์ง๊ธ ์ค์ ๋ก ๋๋ฆฌ๋ ์ค์ |
|
|
| ```bash |
| SRA_EDGE_FIX=1 SRA_RES_CAP_REL=0.03 \ |
| CUDA_VISIBLE_DEVICES=1 python fm_nba_graph_v6.py \ |
| --cfg cfg/nba/cor_fm.yml --exp v3_rel003 \ |
| --batch_size 192 --epochs 150 --fm_in_scaling --tied_noise \ |
| --top_n_neighbors 5 --uncertainty_weight 0.01 --data_dir ./data/nba |
| ``` |
|
|
| eval ๋ณ min-ADEโโ @4.0s (24ํ์ฐจ, 4170/25500 = 16 % ์งํ): |
|
|
| | ์ค์ | 1โ6 | 7โ12 | 13โ18 | 19โ24 | |
| |---|---|---|---|---| |
| | **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 0.828 0.832 0.818 **0.816** 0.819 | **0.803** 0.823 0.808 0.818 0.807 0.820 | |
| | 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 0.863 0.893 0.873 0.871 0.874 | 0.869 0.915 0.873 0.867 **0.830** 0.894 | |
| | 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 โ ๋ฐ์ฐ | ์ข
๋ฃ (์ต๊ทผ 1.41) | |
| | 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.41) | |
| | cap ์์ | 1.141 1.052 1.048 **1.828** 1.429 1.716 | ๋ฐ์ฐ | โ | โ | |
|
|
| ํ์ฌ best (ADE/FDE ๋ **๊ฐ์ ํ๊ฐ ์์ **์์ ์ง์ง์): |
|
|
| | ์ค์ | best ADE / FDE | @eval | ํ๊ฐ ํ์ | |
| |---|---|---|---| |
| | **relcap 0.03** | **0.8029 / 0.9635** | 19 | 24 (์์ ) | |
| | relcap 0.10 | 0.8302 / 1.0857 | 23 | 24 (์ง๋) | |
|
|
| --- |
|
|
| ## 5. ์ด ๋ณ๊ฒฝ์ด ๋
ผ๋ฌธ ์์น์ ์๋ฏธํ๋ ๊ฒ โ ์ ์งํ ์ํ |
|
|
| **โ ๋
ผ๋ฌธ์ MoFlow+SRA = 0.695 ๋ ๋ฒ๊ทธ ์ํ์ ๊ฐ์ด๋ค.** ์ฌ์ ๊ฐ๋ก์ง๋ฅด๋ ๊ทธ๋ํ๋ก |
| ์ป์ ๊ฐ์ด๋ฏ๋ก, "์์ธก๋ ๋ฏธ๋ ์์ ์ฌ ๋ด๋ถ sparse graph ๋๋ถ"์ด๋ผ๋ ๋
ผ๋ฌธ์ ์ค๋ช
๊ณผ |
| ๊ทธ ์์น๋ฅผ ๋ง๋ ์ฝ๋๊ฐ ์ผ์นํ์ง ์๋๋ค. |
|
|
| **โก ์์ ํ ์ฌํ์ต์ ์์ง 0.695 ์ ๋๋ฌํ์ง ๋ชปํ๋ค.** ํ์ฌ 0.8029 (16 % ์งํ). |
| ๋จ์ 84 % ์ cosine LR ๊ฐ์ ์์ ๋ ๋ด๋ ค๊ฐ์ผ ํ๋ฉฐ, **๋๋ฌํ์ง ๋ชปํ ๊ฐ๋ฅ์ฑ์ |
| ์ด๋ ค ์๋ค.** |
|
|
| **โข `SRA_RES_CAP_REL` ์ ๋
ผ๋ฌธ์ ์๋ ํญ์ด๊ณ ์ถ๋ก ์์๋ ์ ์ฉ๋๋ค.** |
| ๋ฐ๋ผ์ ์ด๊ฑด ํ์ดํผํ๋ผ๋ฏธํฐ๊ฐ ์๋๋ผ **๋ฉ์๋ ๋ณ๊ฒฝ**์ด๋ค. ์ด๋ค๋ฉด ๋
ผ๋ฌธ ๋ณธ๋ฌธ์ |
| ๊ธฐ์ ํด์ผ ํ๊ณ , ์ ์ฐ๋ฉด edge ์์ ๋ณธ MoFlow ๋ ์์ ํ์ต๋์ง ์๋๋ค |
| (`SRA_EDGE_FIX=1` ๋ง์ผ๋ก๋ 4ํ์ฐจ์ 1.828 ๋ก ๋ฐ์ฐ). |
| |
| **โฃ ๋ค๋ฅธ ํธ์คํธ๋ ์ด ์ฒ๋ฐฉ์ ์ฐ์ง ์๋๋ค.** MID ๋ `graph_gate_init`/`warmup`, |
| LED ๋ ์ฒ๋ฐฉ ์์ด ํ์ต๋๋ค. ์๋ cap ์ ๋ค๋ฅธ ํธ์คํธ์ ์ ์ฉํ๋ ค๋ฉด ๊ทธ ํธ์คํธ์ |
| residual/embedding ๋น์จ์ ๋จผ์ ์ธก์ ํด์ผ ํ๋ค. |
| |
| **โค ๋น์จ๊ฐ ํ๋์ด ๋๋์ง ์์๋ค.** 0.03 ์ด 0.10 ๋ณด๋ค ๋ซ๋ค. ๋ ์กฐ์ธ ๊ฐ(0.01)์ด |
| ๋์์ง, 0.03 ์ด ์ด๋ฏธ ๊ณผํ ์ ์ฝ์ธ์ง๋ ๋ฏธ๊ฒ์ฆ์ด๋ค. |
| |