po03087's picture
Fix edge-index scene mixing; add relative residual cap; guard LED sigma NaN
37c61d4 verified
|
Raw
History Blame Contribute Delete
8.45 kB

SRA_RES_CAP_REL โ€” ์ƒ๋Œ€ residual cap

ํ•œ ์ค„ ์š”์•ฝ edge ๋ฒ„๊ทธ๋ฅผ ๊ณ ์น˜์ž MoFlow ์ƒ˜ํ”Œ๋ง์ด ๋ฐœ์‚ฐํ–ˆ๋‹ค. ์ ˆ๋Œ€ cap(SRA_RES_CAP)์€ ๊ฐ’์„ ๋‚ฎ์ถฐ๋„ ๋ฐœ์‚ฐ ์‹œ์ ์„ ๋ฏธ๋ฃฐ ๋ฟ์ด์—ˆ๋‹ค. ์ƒํ•œ์„ ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ norm ์— ๋น„๋ก€ํ•˜๋„๋ก ๋ฐ”๊พธ์ž (โ€–resโ€– โ‰ค ratioยทโ€–origโ€–) 13ํšŒ ํ‰๊ฐ€๊นŒ์ง€ ์•ˆ์ •์ ์œผ๋กœ ํ•˜๊ฐ•ํ•˜๋ฉฐ baseline ์„ ์•ž์„ฐ๋‹ค.


1. ์™œ ์ ˆ๋Œ€ cap ์œผ๋กœ๋Š” ๋ถ€์กฑํ•œ๊ฐ€

๋ฐœ์‚ฐ์˜ ์›์ธ์€ MoFlow ์˜ train/sample mismatch ๋‹ค. ํ•™์Šต์€ ๋žœ๋ค timestep ํ•˜๋‚˜์—์„œ ๊ทธ๋ž˜ํ”„๋ฅผ ํ•œ ๋ฒˆ๋งŒ ์ ์šฉํ•˜์ง€๋งŒ, ์ƒ˜ํ”Œ๋ง์€ 10 ์Šคํ… flow ์ ๋ถ„์—์„œ ๋งค ์Šคํ… ์ ์šฉํ•˜๊ณ  ๊ทธ ์ถœ๋ ฅ์ด ๋‹ค์Œ ์Šคํ… ์ž…๋ ฅ์œผ๋กœ ๋˜๋จน์ž„๋œ๋‹ค โ†’ ์„ญ๋™์ด ๋ˆ„์ ๋œ๋‹ค.

์ ˆ๋Œ€ cap ์€ โ€–resโ€– โ‰ค c ๋กœ ๊ณ ์ • ํฌ๊ธฐ๋ฅผ ๊ฐ•์ œํ•œ๋‹ค. ๋ฌธ์ œ๋Š” ๋‘ ๊ฐ€์ง€๋‹ค.

  1. ํ˜ธ์ŠคํŠธ ์Šค์ผ€์ผ ์˜์กด. residual norm ์€ ๊ทธ ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ์ด ์“ฐ๋Š” ์Šค์ผ€์ผ ์œ„์— ์žˆ๋‹ค. ์ธก์ •๊ฐ’: MoFlow ์—์„œ ๊ทธ๋ž˜ํ”„ residual โ‰ˆ 6.3, ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ โ‰ˆ 16. MoFlow ์—์„œ ํŠœ๋‹ํ•œ c = 3.0 ์ด ์ž„๋ฒ ๋”ฉ ์Šค์ผ€์ผ์ด ๋‹ค๋ฅธ MID/LED ์—์„œ๋Š” ์‚ฌ์‹ค์ƒ ๋ฌด์—ฐ์‚ฐ์ด๊ฑฐ๋‚˜ ๋ฐ˜๋Œ€๋กœ ๊ณผ๋„ํ•œ ์ œ์•ฝ์ด ๋œ๋‹ค.
  2. ๋ˆ„์ ์„ ๊ฒฐ์ •ํ•˜๋Š” ๊ฒƒ์€ ์ ˆ๋Œ€ ํฌ๊ธฐ๊ฐ€ ์•„๋‹ˆ๋ผ ๋น„์œจ. ์Šคํ…๋‹น ์ƒํƒœ๊ฐ€ x โ† x + res ๋กœ ๊ฐฑ์‹ ๋  ๋•Œ ๋ฐœ์‚ฐ ์—ฌ๋ถ€๋ฅผ ์ง€๋ฐฐํ•˜๋Š” ๊ฒƒ์€ โ€–resโ€–/โ€–xโ€– ๋‹ค. ์ ˆ๋Œ€ cap ์€ ์ด ๋น„์œจ์„ ์ง์ ‘ ํ†ต์ œํ•˜์ง€ ๋ชปํ•œ๋‹ค.

์‹ค์ œ๋กœ ์ ˆ๋Œ€ cap ์€ ๊ฐ’์„ ๋‚ฎ์ถœ์ˆ˜๋ก ๋ฐœ์‚ฐ์ด ๋ฏธ๋ค„์งˆ ๋ฟ ์‚ฌ๋ผ์ง€์ง€ ์•Š์•˜๋‹ค:

์ ˆ๋Œ€ cap ๋ฐœ์‚ฐ ์‹œ์ 
6.0 eval 4
3.0 eval 5
1.0 eval 9 (์ดํ›„ 1.15~1.61 ์ง„๋™)
0.5 12ํšŒ์ฐจ๋ถ€ํ„ฐ ๋ฐ˜๋“ฑ (0.873 โ†’ 0.922 โ†’ 0.958)

2. ๊ตฌํ˜„

# SRA_RES_CAP_REL: ๋…ธ๋“œ๋ณ„ ์ƒํ•œ์„ ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ norm ์— ๋น„๋ก€ํ•ด ์ •ํ•œ๋‹ค.
_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)        # [N, 1] ๋…ธ๋“œ๋ณ„ ์ƒํ•œ
    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
  • orig = ํ˜ธ์ŠคํŠธ ์ž„๋ฒ ๋”ฉ, res = ๊ทธ๋ž˜ํ”„ ์„ญ๋™. ์ƒํ•œ์ด ๋…ธ๋“œ๋งˆ๋‹ค ์ž๊ธฐ ์ž„๋ฒ ๋”ฉ ํฌ๊ธฐ์— ๋น„๋ก€ํ•ด ์ •ํ•ด์ง€๋ฏ€๋กœ ์Šค์ผ€์ผ ๋ฌด๊ด€ํ•˜๋‹ค.
  • torch.where ๋ฅผ ์“ฐ๋Š” ์ด์œ ๋Š” ยง4 ์ฐธ์กฐ (์ด์ „ ๊ตฌํ˜„์˜ ์น˜๋ช…์  ๋ฒ„๊ทธ).
  • ๊ธฐ๋ณธ๊ฐ’ 0 = ๊บผ์ง. ์ผœ์ง€ ์•Š์œผ๋ฉด ์›๋ณธ๊ณผ ๋™์ผํ•˜๊ฒŒ ๋™์ž‘ํ•œ๋‹ค.

๊ฒ€์ฆ (A=11, K=4):

์„ค์ • out_proj grad โ€–resโ€–max โ€–resโ€–/โ€–origโ€– max
๋ฌด์ œํ•œ 1,716,997 9.009 0.557
์ ˆ๋Œ€ cap 1.0 255,169 1.000 0.071
relcap 0.10 408,742 1.765 0.100 โœ“
relcap 0.03 122,623 0.529 0.030 โœ“

๋น„์œจ์ด ์ง€์ •๊ฐ’์œผ๋กœ ์ •ํ™•ํžˆ ์ œํ•œ๋˜๊ณ  gradient ๋„ ์ •์ƒ์ด๋‹ค.


3. ๊ฒฐ๊ณผ (MoFlow-NBA, full SRA, SRA_EDGE_FIX=1)

eval ๋ณ„ min-ADEโ‚‚โ‚€ @4.0s:

eval 1 2 3 4 5 6 7 8 9 10 11 12 13
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
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
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
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
baseline(์ฐธ๊ณ ) 1.163 1.021 0.961 0.964 0.952 0.955 0.953 0.895 โ€” โ€” โ€” โ€” โ€”

ํ˜„์žฌ best (ADE/FDE ๋Š” ๊ฐ™์€ ํ‰๊ฐ€ ์‹œ์ ์—์„œ ์ง์ง€์Œ):

์„ค์ • best ADE / FDE @eval ์ง„ํ–‰
relcap 0.03 0.8399 / 1.0025 13 2368/25500 (9 %)
relcap 0.10 0.8568 / 1.0441 11 2349/25500
cap 0.5 0.8732 / 1.0531 10 2379/25500
cap 1.0 0.9662 / 1.2508 8 3390/25500 โ€” ๋ฐœ์‚ฐ

๊ด€์ฐฐ:

  • ์›๋ž˜ ๋ฐœ์‚ฐ ์ง€์ (eval 4~5)์„ ์ƒ๋Œ€ cap 2์ข… ๋ชจ๋‘ ํ†ต๊ณผํ–ˆ๋‹ค. ์ ˆ๋Œ€ cap 1.0 ๋„ 8ํšŒ์ฐจ๊นŒ์ง€๋Š” ๋ฉ€์ฉกํ•ด ๋ณด์˜€์œผ๋ฏ€๋กœ 4~5ํšŒ ํ†ต๊ณผ๋งŒ์œผ๋กœ๋Š” ๋ถ€์กฑํ•œ๋ฐ, 13ํšŒ๊นŒ์ง€ ์œ ์ง€๋œ ๊ฒƒ์€ ๋‹ค๋ฅธ ์ˆ˜์ค€์˜ ๊ทผ๊ฑฐ๋‹ค.
  • baseline(0.895) ์„ ์•ž์„ฐ๋‹ค. relcap 0.03 ์ด 0.840 ์œผ๋กœ โˆ’0.055.
  • ์ ˆ๋Œ€ cap 0.5 ๋Š” 10ํšŒ์ฐจ 0.873 ์ดํ›„ 0.922 โ†’ 0.958 ๋กœ 3ํšŒ ์—ฐ์† ์ƒ์Šน โ€” ์ƒ๋Œ€ cap ๊ณผ ๋‹ฌ๋ฆฌ ๋ถˆ์•ˆ์ • ์กฐ์ง์ด ๋ณด์ธ๋‹ค.
  • ๋น„์œจ์„ 0.10 โ†’ 0.03 ์œผ๋กœ ๋” ์กฐ์—ฌ๋„ ์„ฑ๋Šฅ์ด ๋‚˜๋น ์ง€์ง€ ์•Š์•˜๋‹ค. ์ฆ‰ ๋ฌด์ œํ•œ ์ƒํƒœ์˜ ๋น„์œจ 0.557 ์€ ํ•„์š” ์ด์ƒ์œผ๋กœ ํฌ๊ณ , ๊ทธ ๊ผฌ๋ฆฌ๊ฐ€ ๋ฐœ์‚ฐ์„ ์œ ๋ฐœํ–ˆ๋‹ค๋Š” ํ•ด์„๊ณผ ์ผ์น˜ํ•œ๋‹ค.

4. โš ๏ธ ๊ฐ™์ด ๊ณ ์นœ ๊ฒƒ โ€” ์ด์ „ cap ๊ตฌํ˜„์˜ ์น˜๋ช…์  ๋ฒ„๊ทธ

์ฒ˜์Œ ์ž‘์„ฑํ•œ cap ์€ ์ด๋žฌ๋‹ค:

res = 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 ํ•ฉ):

์„ค์ • grad
SOFT_START ๋งŒ 78,751
RES_CAP ๋งŒ 765,506
SOFT_START + RES_CAP 0.00 โ† ํ•™์Šต ๋ถˆ๊ฐ€

์ด ์กฐํ•ฉ์œผ๋กœ ๋Œ๋ฆฐ ์‹คํ–‰๋“ค์˜ ์ฒดํฌํฌ์ธํŠธ๋ฅผ ์—ด์–ด๋ณด๋‹ˆ 58 epoch ๋’ค์—๋„ future_graph.out_proj ๊ฐ€ ์ •ํ™•ํžˆ 0.000e+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

5. ์‚ฌ์šฉ๋ฒ•

# ๊ถŒ์žฅ ์„ค์ • (MoFlow-NBA)
SRA_EDGE_FIX=1 SRA_RES_CAP_REL=0.03 \
CUDA_VISIBLE_DEVICES=0 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
ํ† ๊ธ€ ๊ธฐ๋ณธ ์—ญํ• 
SRA_EDGE_FIX off edge-index ์”ฌ ํ˜ผํ•ฉ ๋ฒ„๊ทธ ์ˆ˜์ • (ํ•„์ˆ˜)
SRA_RES_CAP_REL 0 (off) ๋น„์œจ ์ƒํ•œ โ€–resโ€– โ‰ค rยทโ€–origโ€– โ€” ๊ถŒ์žฅ
SRA_RES_CAP 0 (off) ์ ˆ๋Œ€ ์ƒํ•œ โ€” MoFlow ์—์„œ ์—ด๋“ฑํ•จ์ด ํ™•์ธ๋จ
SRA_SOFT_START off RES_CAP ๊ณผ ํ•จ๊ป˜ ์“ฐ์ง€ ๋ง ๊ฒƒ (ยง4). ๋‹จ๋…์œผ๋กœ๋„ ๋ฐœ์‚ฐ ๋ชป ๋ง‰์Œ
SRA_GATE_SCALE 1.0 ์ „์—ญ ์ถ•์†Œ โ€” ๊ฐ€์žฅ ๋นจ๋ฆฌ ๋ฐœ์‚ฐ(3ํšŒ์ฐจ), ํ๊ธฐ

6. ๋ฏธํ™•์ • / ํ•œ๊ณ„

โ‘  ์ตœ์ข… ์ˆ˜๋ ด๊ฐ’์€ ์•„์ง ๋ชจ๋ฅธ๋‹ค. ํ˜„์žฌ 9 %(2368/25500) ์—์„œ 0.840 ์ด๋‹ค. ๊ธฐ์กด(๋ฒ„๊ทธํŒ) SRA = 0.695, baseline = 0.703. ๋‚จ์€ 91 % ์™€ cosine LR ๊ฐ์‡ ์—์„œ ๋” ๋‚ด๋ ค๊ฐ€์•ผ ํ•˜๋ฉฐ, 0.695 ์— ๋„๋‹ฌํ•˜์ง€ ๋ชปํ•  ๊ฐ€๋Šฅ์„ฑ์€ ์—ฌ์ „ํžˆ ์—ด๋ ค ์žˆ๋‹ค.

โ‘ก ๋น„์œจ๊ฐ’ ํŠœ๋‹์ด ๋๋‚˜์ง€ ์•Š์•˜๋‹ค. 0.03 ๊ณผ 0.10 ์ด ๋น„์Šทํ•˜๊ณ  0.03 ์ด ๊ทผ์†Œ ์šฐ์œ„๋‹ค. ๋” ์กฐ์ธ ๊ฐ’(0.01)์ด ๋‚˜์„์ง€, ์•„๋‹ˆ๋ฉด 0.03 ์ด ์ด๋ฏธ ๊ณผ๋„ํ•œ ์ œ์•ฝ์ธ์ง€๋Š” ๋ฏธ๊ฒ€์ฆ์ด๋‹ค. cap ์„ ์กฐ์ผ์ˆ˜๋ก ๋ฐœ์‚ฐ์€ ๋ง‰ํžˆ์ง€๋งŒ ๊ทธ๋ž˜ํ”„ ๊ธฐ์—ฌ๋„ ์ค„์–ด๋“œ๋ฏ€๋กœ, "๋ฐœ์‚ฐ ์•ˆ ํ•˜๋ฉด์„œ baseline ๋ณด๋‹ค ๋‚˜์€" ๊ตฌ๊ฐ„์ด ์‹ค์ œ๋กœ ์กด์žฌํ•˜๋Š”์ง€๊ฐ€ ์ตœ์ข… ์งˆ๋ฌธ์ด๋‹ค. ํ˜„์žฌ 0.840 vs baseline 0.895 ๋Š” ๊ทธ ๊ตฌ๊ฐ„์ด ์กด์žฌํ•œ๋‹ค๋Š” ์ฒซ ์ฆ๊ฑฐ๋‹ค.

โ‘ข MoFlow ์ „์šฉ ์ฒ˜๋ฐฉ์ด๋‹ค. MIDยทLED ๋Š” iterative denoiser ๋ผ ์ด ๋ฐœ์‚ฐ์ด ์—†๊ณ , RES_CAP* ์„ ์“ฐ์ง€ ์•Š๋Š”๋‹ค (MID ๋Š” gate_init/warmup, LED ๋Š” ์ฒ˜๋ฐฉ ์—†์Œ). ๋‹ค๋ฅธ ํ˜ธ์ŠคํŠธ์— ์ ์šฉํ•˜๋ ค๋ฉด ๊ทธ ํ˜ธ์ŠคํŠธ์˜ residual/embedding ๋น„์œจ์„ ๋จผ์ € ์ธก์ •ํ•ด์•ผ ํ•œ๋‹ค.

โ‘ฃ ์ ˆ๋Œ€ cap ์ด ํ•ญ์ƒ ๋‚˜์˜๋‹ค๊ณ  ๋‹จ์ •ํ•  ์ˆ˜๋Š” ์—†๋‹ค. cap 0.5 ๋Š” ์•„์ง ๋ฐœ์‚ฐํ•˜์ง€ ์•Š์•˜๊ณ  ๋ฐ˜๋“ฑ ์กฐ์ง๋งŒ ๋ณด์ธ๋‹ค. ์ƒ๋Œ€ cap ์˜ ์šฐ์œ„๋Š” ํ˜„์žฌ 13ํšŒ ํ‰๊ฐ€ ๊ธฐ์ค€์˜ ๊ด€์ฐฐ์ด๋ฉฐ, ์™„์ฃผ๊นŒ์ง€ ๊ฐ€์•ผ ํ™•์ •๋œ๋‹ค.