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 ์ฌ ํผํฉ ๋ฒ๊ทธ
๋ฐ๋ ์ฝ๋
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 โ ์ด๊ธฐ ์ถฉ๊ฒฉ ์ ๊ฑฐ
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 โ ์ ์ญ ์ถ์
res = SRA_GATE_SCALE * gate * self.out_proj(nodes) # ๊ธฐ๋ณธ 1.0 = off
๊ฐ์ฅ ๋จ์ํ ์ฒ๋ฐฉ์ด์ง๋ง ๊ฐ์ฅ ๋นจ๋ฆฌ ๋ฐ์ฐํ๋ค(3ํ์ฐจ). ํ๊ธฐ.
2.3 SRA_RES_CAP / SRA_RES_CAP_REL โ residual norm ์ํ
_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)) ์ด์๋๋ฐ ๋ ๊ฐ์ง๊ฐ ํ๋ ธ๋ค.
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๊ฑด์ ํ๊ธฐํ๋ค.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. ์ง๊ธ ์ค์ ๋ก ๋๋ฆฌ๋ ์ค์
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 ์ด ์ด๋ฏธ ๊ณผํ ์ ์ฝ์ธ์ง๋ ๋ฏธ๊ฒ์ฆ์ด๋ค.