Skip to content

Fix: Corrected the content related to issue #18 and added a test to verify its correctness. - #19

Open
Antikira wants to merge 12 commits into
toruseo:mainfrom
Antikira:dev/#18_fix
Open

Antikira wants to merge 12 commits into
toruseo:mainfrom
Antikira:dev/#18_fix

Conversation

@Antikira

Copy link
Copy Markdown

Fixes #18

Summary

Fixed an issue where JAX node model ignored discharge from the 3rd and later inlinks at merge nodes with 3+ inlinks.

Changes

1. unsim/unsim_diff.py

Vectorized the multi-inlink merge priority logic from Node._transfer_merge (unsim/unsim.py) into a vectorize:

  1. Distribute supply based on merge priorities ($\alpha_i$) and link demands.
  2. Use cumulative sumation over unsatisfied demands to sequentially allocate remaining supply.
  3. Unified all merge nodes into this vectorized formula.

2. tests/test_diff.py

  • Added test_merge_3inlinks to verify that Python and JAX simulation total travel time match exactly on a 3-inlink congested merge network.

@toruseo

toruseo commented Sep 12, 2026

Copy link
Copy Markdown
Owner

@Antikira
ありがとうございます.

Issueのほうに書いてある通り,この修正にはdesign decisionが必要です.
今のAIコーディング時代には人間の理解と意思決定が重要になる(まあ一時的なものかもしれませんが)ので,Antikiraさんご自身の判断についてご説明をお願いします.

あと,本リポジトリのライセンスをMITからApache 2.0に変更します.Antikira さんの貢献もApacheになることになりますが,よろしいでしょうか?

@Antikira

Copy link
Copy Markdown
Author

@toruseo
レビューありがとうございます.

今回の実装方針を採用した理由は主に以下の2点です.

  1. 既存のPython実装との整合性
    後の検証や挙動の確認への優位性を期待し,Python実装と同等の手法で合わせることを優先しました.

  2. 実行速度の優位性
    INMにマッピングする方式と比較したところ,わずかではありますが高速化が確認できました.

    • 検証条件: 3リンクが1リンクに合流するシンプルな合流部ネットワーク
    • 結果: 3回の平均で約1.1倍の高速化を確認(※厳密な統計的有意性を示すものではなく,簡易的なベンチマーク結果です)

割り当て規則の同等性を保ちつつ,速度面でもメリットが見込めるため,今回はこちらの実装方式を選択しています.

ライセンス変更の件は同意いたします.
よろしくお願いします.

@toruseo

toruseo commented Sep 12, 2026

Copy link
Copy Markdown
Owner

ありがとうございます.自分もその設計方針に賛成です.

astra先生にレビューしてもらったら以下のコメントをもらいました.妥当に思えるので,修正お願いします.


対象はコミット e66a512 です。マージ前に修正が必要です。総流量の勾配が誤るケースと、小さい正の優先度でPython側と流量が一致しないケースを再現しました。 通常の優先度での流量計算は、Python側の配分に対応しています。(GitHub)

1. 【P1】優先配分量と需要が一致する点で、総流量の勾配が誤る

箇所:unsim/unsim_diff.py の base_q、rem_S、extra_q の計算。

このPRは、既存の2本合流の中央値式も、3本以上に対応する式へ置き換えています。置き換え後の式で、流量は一致しても勾配が誤るケースがあります。(GitHub)

例えば、2本の流入リンクについて、需要を D = [0.5, 1.0]、優先度を p = [1, 1]、供給を S = 1.0 とします。流量は両実装とも [0.5, 0.5] です。

総流量を (Q=q_1+q_2) とすると、この近傍では供給制約が有効で、(Q=S) が成立します。該当計算式を単体実行した結果は以下です。

微分する量 正しい値 修正前のJAX式 このPR
需要 (D_1) に対する総流量の微分 0 0 0.3125
優先度 (p_1) に対する総流量の微分 0 0 −0.078125
供給 (S) に対する総流量の微分 1 1 0.84375

float32・float64、JITあり・なしで同じ結果を再現しました。3本合流でも、D = [0.25, 0.75, 0.75]、p = [1, 1, 2]、S = 1.0 とすると、需要 D[0] に対する総流量の微分が 0.3125 になります。

この例では、base_q、rem_S、extra_q に含まれる複数の minimum・maximum が同時に等値点に入り、それぞれの微分を合成した結果、総流量の感度が崩れています。総流量自体は、この近傍で (Q=S) と微分可能です。 下流リンクへの流入量に、需要や優先度に対する誤った感度が入ります。(GitHub)

修正要求: 2本合流は既存の中央値式を保持するのが最小変更です。3本以上の合流については、Python側の配分を保ちつつ、供給制約下で総流量の微分が「需要について0、優先度について0、供給について1」になるよう計算式を修正する必要があります。

2. 【P2】正規化分母の下限により、Python側と配分が変わる

箇所:alphas = p / jnp.maximum(jnp.sum(p, ...), 1e-10)。

Python側は優先度の総和で正規化しています。このPRでは総和に 1e-10 の下限を設けているため、総和がそれを下回ると配分が変わります。

例えば、D = [1, 1, 1]、S = 1、p = [1e-12, 2e-12, 1e-12] の場合、単体実行の結果は次のとおりです。

Python側: [0.25, 0.50, 0.25]
このPR:   [0.97, 0.02, 0.01]  (丸めて表示)

優先度を同じ比率の [1, 2, 1] にすると、このPRも [0.25, 0.50, 0.25] を返します。したがって、優先度の絶対的な大きさによって配分が変わっています。

修正要求: 正の総優先度には、その値を分母として使う必要があります。パディングによって総優先度がゼロになる行は、分母を別途安全化できます。

3. 追加テストは、リンク別の挙動と勾配まで検証する必要がある

追加された test_merge_3inlinks が比較しているのは、総旅行時間だけです。また、使用している equal_tolerance のデフォルトは 相対許容誤差10% です。

総旅行時間という集計値だけでは、リンク間の配分の違いを捉えきれません。また、この追加テストには自動微分の検査が含まれていません。

ご指定の設計判断に対応するテストとして、次の3点が必要です。

  • リンク別の一致: 各リンクの累積流入・累積流出を時刻ごとに比較する。需要不足のリンクから余剰供給を再配分するケースと、4本以上の合流も含める。
  • 勾配の一致: 滑らかな点で、需要・供給・優先度に対する自動微分を有限差分と比較する。優先配分量と需要が一致する点では、上記の総流量の微分を検査する。
  • 優先度の倍率に対する不変性: 優先度全体を同じ正数倍しても、流量が変わらないことを検査する。

設計方針についての確認結果

Python側の「優先度による初期配分と、登録順での余剰供給の再配分」をJAX側へ移植する方針は、挙動統一という目的に対応しています。 world_to_jax はPython側と同じ inlinks.values() の順番を保持しており、再配分の順序も対応しています。

今回、優先度を0.1〜4.0として、2・3・4・8本合流を各2,000ケース、計8,000ケース比較しました。リンク別流量の最大絶対差は約 (5.7\times10^{-7}) でした。さらに、滑らかな点400ケースでは、需要・供給・優先度に対する自動微分とPython側の有限差分が、最大絶対差約 (2.1\times10^{-10}) で一致しました。

したがって、現在の配分方針を維持し、等値点での勾配と優先度の正規化を修正するのが、この設計判断に沿った対応です。

検証は合流計算を抽出した単体実行です(JAX 0.9.0.1、CPU)。リポジトリ全体の pytest は未実行です。検証スクリプトと実行結果

@Antikira

Copy link
Copy Markdown
Author

修正しました.
修正内容とご指摘内容の対応は以下のとおりです.

修正内容

P1関連: 3cdfabc

最大化と最小化関数が入力値が等しい場合に微分を0.5に割り当てる仕様のため発生していました.
検証コードはappendix参照.
勾配の均等割り当てにより発生する問題を計算グラフの構造上キャンセルできるように余計な計算を追加しました.
理論的背景は以下のとおりです.
まず,バグの発生原因としてはもとのリンク流量 $\partial q^{\text{raw}}_i$ のある変数 $x$ による偏微分 $\frac{\partial q^{\text{raw}}_i}{\partial x}$ が正常に計算できないことです.
そこで $q_i = f(q^{\text{raw}}_i)$としたときに $\frac{\partial q_i}{\partial x}$ が好ましい微分値になるような関数 $f$ を検討したところ,以下の処理が該当しそうであると確認しました.

$$q_i = q_i^{\text{raw}} \cdot \frac{S}{\sum_j q_j^{\text{raw}}}$$

詳細な検討はAppendixのとおりです.

P2関連: b845183

もとの実装は0除算エラー対策です.
そのためwhere関数で処理しました.

P3 関連: a39da64, a46d993, cc7af01, 3390d9e

  1. 流量別の検討(a46d993)を追加しました.
    検討手法はリンク別,自己区別に計算結果を個別に検討しています.
  • 実装クラス: TestNumericalAgreement
  • 追加テスト関数:
    • test_merge_2inlinks_fair_linkwise: 2流入リンク・均等需要でのリンク別流量検証
    • test_merge_2inlinks_unfair_linkwise: 2流入リンク・不均等需要でのリンク別流量検証
    • test_merge_3inlinks_linkwise: 3流入リンクでのリンク別流量検証
    • test_merge_surplus_reallocation_linkwise: 余剰容量の再配分(Surplus Reallocation)検証
    • test_merge_4inlinks_linkwise: 4流入リンクでのリンク別流量検証
    • test_priority_scale_invariance_formula: 優先度定数倍に対する計算式のスケール不変性検証
    • test_tiny_priority_allocation: 極小優先度時の配分挙動検証
    • test_simulation_scale_invariance: 優先度定数倍に対するシミュレーション全体のスケール不変性検証
  1. 勾配検証(cc7af01, 3390d9e)を追加しました.
    事前の実装と現在の実装,理論値が同一性の検証を追加しました.
  • 実装クラス: TestGradient
  • 追加テスト関数:
    • test_grad_merge_2inlinks_vs_prior_smooth: スムーズな混雑領域における事前実装との勾配比較
    • test_grad_merge_2inlinks_prior_at_boundary: 境界点における勾配挙動の検証
    • test_grad_merge_2inlinks_vs_prior_uncongested: 非混雑レジームでの勾配検証
    • test_grad_merge_2inlinks_vs_prior_one_under_capacity: 1リンクのみ容量以下のケースでの勾配検証
    • test_grad_merge_2inlinks_boundary: 2流入リンク境界値における勾配検証($\partial Q/\partial D=[0, 0], \partial Q/\partial S=1.0$)
    • test_grad_merge_3inlinks_boundary: 3流入リンク境界値における勾配検証

Appendix

jaxの勾配に関して

以下のコードで勾配を追跡したところ局地処理が問題の原因だと特定しました.

# 誤差の検証
# maxとminの微分のせいであることを確認する.
# 参考: https://github.com/jax-ml/jax/issues/13600

#%%
import jax.numpy as jnp
import jax
#%%
D1 = 0.5
D2 = 1.0
p1 = 1.0
p2 = 1.0
S = 1.0
def bug_func(D1, D2, p1, p2, S):
    alpha1 = p1 / (p1 + p2)
    alpha2 = p2 / (p1 + p2)
    q1 = jnp.minimum(D1, alpha1 * S)
    q2 = jnp.minimum(D2, alpha2 * S)
    rem_S1 = jnp.maximum(S - (q1 + q2), 0.0)
    e1 = jnp.minimum(D1 - q1, rem_S1)
    q1 = q1 + e1
    
    rem_S2 = jnp.maximum(S - (q1 + q2), 0.0)
    e2 = jnp.minimum(D2 - q2, rem_S2)
    q2 = q2 + e2
    
    return q1 + q2
def differentiable_mid_2in(D1, D2, p1, p2, S):
    alpha1 = p1 / (p1 + p2)
    alpha2 = p2 / (p1 + p2)
    q1 = jnp.minimum(D1, jnp.maximum(S - D2, alpha1 * S))
    q2 = jnp.minimum(D2, jnp.maximum(S - D1, alpha2 * S))
    
    return q1+q2
# %%
original_grad_D1 = jax.grad(differentiable_mid_2in, argnums=0)
original_grad_p1 = jax.grad(differentiable_mid_2in, argnums=2)
original_grad_S  = jax.grad(differentiable_mid_2in, argnums=4)
bug_grad_D1 = jax.grad(bug_func, argnums=0)
bug_grad_p1 = jax.grad(bug_func, argnums=2)
bug_grad_S  = jax.grad(bug_func, argnums=4)

print("もとの実装")
print(f"dQ/dD1 : {original_grad_D1(D1, D2, p1, p2, S):.6f} (理論値 0.0)")
print(f"dQ/dp1 : {original_grad_p1(D1, D2, p1, p2, S):.6f} (理論値 0.0)")
print(f"dQ/dS  : {original_grad_S(D1, D2, p1, p2, S):.6f} (理論値 1.0)")

print("バグ再現")
print(f"dQ/dD1 : {bug_grad_D1(D1, D2, p1, p2, S):.6f} (理論値 0.0)")
print(f"dQ/dp1 : {bug_grad_p1(D1, D2, p1, p2, S):.6f} (理論値 0.0)")
print(f"dQ/dS  : {bug_grad_S(D1, D2, p1, p2, S):.6f} (理論値 1.0)")
#%%

print("\n" + "=" * 85)
print("連鎖律")
print("=" * 85)

# 全体の中間変数を一括追跡する関数
def bug_trace_all(D1, D2, p1, p2, S):
    alpha1 = p1 / (p1 + p2)
    alpha2 = p2 / (p1 + p2)
    a1_S = alpha1 * S
    a2_S = alpha2 * S
    q1_0 = jnp.minimum(D1, a1_S)
    q2_0 = jnp.minimum(D2, a2_S)
    diff_rem1 = S - (q1_0 + q2_0)
    rem_S1 = jnp.maximum(diff_rem1, 0.0)
    diff_D1 = D1 - q1_0
    e1 = jnp.minimum(diff_D1, rem_S1)
    q1_1 = q1_0 + e1
    
    diff_rem2 = S - (q1_1 + q2_0)
    rem_S2 = jnp.maximum(diff_rem2, 0.0)
    diff_D2 = D2 - q2_0
    e2 = jnp.minimum(diff_D2, rem_S2)
    q2_1 = q2_0 + e2
    Q = q1_1 + q2_1
    
    return [
        ("alpha1",   "p1 / (p1 + p2)",                  alpha1),
        ("alpha2",   "p2 / (p1 + p2)",                  alpha2),
        ("a1_S",     "alpha1 * S",                      a1_S  ),
        ("a2_S",     "alpha2 * S",                      a2_S  ),
        ("q1_0",     "min(D1, a1_S)",                   q1_0  ),
        ("q2_0",     "min(D2, a2_S)",                   q2_0  ),
        ("rem_S1",   "max(S - (q1_0 + q2_0), 0.0)",     rem_S1),
        ("e1",       "min(D1 - q1_0, rem_S1)",          e1    ),
        ("q1_1",     "q1_0 + e1",                       q1_1  ),
        ("rem_S2",   "max(S - (q1_1 + q2_0), 0.0)",     rem_S2),
        ("e2",       "min(D2 - q2_0, rem_S2)",          e2    ),
        ("q2_1",     "q2_0 + e2",                       q2_1  ),
        ("Q",        "q1_1 + q2_1",                     Q     ),
    ]

trace_outputs = bug_trace_all(D1, D2, p1, p2, S)
inputs = (D1, D2, p1, p2, S)

for idx, (var_name, expr, val) in enumerate(trace_outputs):
    fn = lambda d1, d2, p1, p2, s: bug_trace_all(d1, d2, p1, p2, s)[idx][2]
    dD1 = jax.grad(fn, 0)(*inputs)
    dp1 = jax.grad(fn, 2)(*inputs)
    dS  = jax.grad(fn, 4)(*inputs)
    
    print(f"Step {idx+1:2d}: {var_name:8s} = {expr:30s} | 値: {float(val):.4f}")
    print(f"         d{var_name}/dD1 = {dD1:.6f} | d{var_name}/dp1 = {dp1:.6f} | d{var_name}/dS = {dS:.6f}")

証明

$\sum_j q_j^{\text{raw}} = Q^{\text{raw}}$ とすると求めたい個別流量の微分は,

$$\frac{\partial q_i}{\partial x} = \frac{S}{Q^{\text{raw}}} \frac{\partial q_i^{\text{raw}}}{\partial x} - \frac{q_i^{\text{raw}} \cdot S}{(Q^{\text{raw}})^2} \frac{\partial Q^{\text{raw}}}{\partial x}$$

となる.

渋滞中は $Q^{\text{raw}} = S$ が成立するため、代入して整理すると

$$\frac{\partial q_i}{\partial x} = 1 \cdot \frac{\partial q_i^{\text{raw}}}{\partial x} - \frac{q_i^{\text{raw}}}{S} \frac{\partial Q^{\text{raw}}}{\partial x}$$

となる.この総和が個別の流量の和が総流量であることを念頭に置くと線形性から以下が成立する.

$$\frac{\partial Q}{\partial x} = \sum_i \frac{\partial q_i^{\text{raw}}}{\partial x} - \frac{\partial Q^{\text{raw}}}{\partial x} \sum_i \frac{q_i^{\text{raw}}}{S}$$

ここで, $Q^{\text{raw}}$ の定義を鑑みると常に0になることが理解できる.

@toruseo

toruseo commented Sep 13, 2026 •

Copy link
Copy Markdown
Owner

ありがとうございます.

テストコードが乱立していて旧記述なども紛れていているので,いったん整理していただけますか?
さすがにこのためだけに1000行追加するわけにはいきません...

あと,例によってastraが以下の問題を見つけてます:

1.需要合計と供給が等しい場合、上下流の勾配が不整合になる【P1】
該当箇所:unsim/unsim_diff.py の compute_node_transfers()、merge_q_all と merge_q_out の計算。

2.供給が0の場合、再正規化によって上流リンクの勾配が消える【P1】
該当箇所:同関数の scale_cong の計算と、q_cong * scale_cong。

追記:全体的方針はよいと思います

@Antikira Antikira left a comment •

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

修正しました.
テストコードの肥大化に関しては間違えてlintをかけていたことが原因です.
今回共通化できそうな処理を共通化しました.
これ以上の整理が必要な場合は別のPRで行ったほうが紛れがないかと思います.

また実装の修正をおこなしました.
供給0を実現するために後進波速度の計算等一部零除算対策を導入しました.
また,修正項を足し算で加えてみました.妥当性に関しては以下のとおりなので問題ないかと思います.
よろしくお願いします.

Appendix

定式化

合流ノードにおける入力を以下とする:

  • 有効流入リンク数 $N$,各リンクの需要 $D_i \ge 0$
  • 下流供給量 $S_{\text{merge}} \ge 0$
  • 合流優先度 $p_i \ge 0$,正規化優先度 $\alpha_i = \frac{p_i}{\sum_k p_k}$($\sum_{i=1}^N \alpha_i = 1.0$)
  • 合計需要 $\text{total D} = \sum_{i=1}^N D_i$

まず基礎的な交通量$q_{\text{base}, i}$は
$$q_{\text{base}, i} = \min(D_i, \alpha_i S_{\text{merge}})$$
となる.
このときの未配分余剰供給 $\text{rem S}$ は
$$\text{rem S} = S_{\text{merge}} - \sum_{i=1}^N q_{\text{base}, i}$$
である.

各リンクの未充足需要$\text{cap}i$は
$$\text{cap}i = D_i - q{\text{base}, i}$$
である.
この未充足需要に対して,余剰配分供給$q
{\text{extra}, i}$を順次配分すると,
$$\text{cum cap}i = \sum{k=1}^i \text{cap}k$$
$$q
{\text{extra}, i} = \min\left( \text{cap}i, \max\left( \text{rem S} - (\text{cum_cap}i - \text{cap}i), 0 \right) \right)$$
となる.
よって,さしあたりの総流量$q
{\text{raw}, i}$は
$$q
{\text{raw}, i} = q
{\text{base}, i} + q_{\text{extra}, i}$$
である.

このとき,この計算中の最小化,最大化処理により適切な微分が得られないため,何らかの余計な計算を挟んで無理やり勾配を補正する必要がある.

なんらかのいい感じの関数

実際の渋滞時流量$q_{\text{cong}, i}$が好ましい形で得られるように,$$q_{\text{raw}, i}$を処理する関数を検討する.
このとき,$q_{\text{cong}, i}$が満たすべき条件は以下の通りである:

  1. 流量がもとの値と変わらないこと:$q_{\text{cong}, i} = q_{\text{raw}, i}$
  2. 各リンクの需要が供給の優先配分を超える場合において,流量の供給微分$\partial q_{\text{cong}, i} / \partial S_{\text{merge}}=\alpha_i$かつ需要微分$\partial q_{\text{cong}, i} / \partial D_j=0$が成立すること.
  3. 供給が0の場合においては供給による微分が1であり,かつ需要による微分が0であること.

色々探したところ以下の式が良さそうである:
$$q_{\text{cong}, i} = q_{\text{raw}, i} + \alpha_i \left( S_{\text{merge}} - \sum_{k=1}^N q_{\text{raw}, k} \right)$$

渋滞時において,$\sum_{k=1}^N q_{\text{raw}, k} = S_{\text{merge}}$ が成立するため,この補正項は0となる.

自由流時は,どうせ,流量は需要に従うため細かいことは考えなくて良い.
すなわち,実際の流量$q_{\text{all}, i}$は以下になる:
$$q_{\text{all}, i} = \begin{cases} D_i & (\text{total_D} \le S_{\text{merge}}) \ q_{\text{cong}, i} & (\text{total_D} > S_{\text{merge}}) \end{cases}$$
$$q_{\text{out}} = \sum_{i=1}^N q_{\text{all}, i}$$

微分の確認

完全に渋滞でない場合

完全渋滞ではない場合の流量は$q_{\text{all}, i} = D_i$であり,このときの微分は$\partial q_{\text{all}, i} / \partial D_j = \delta_{ij}$,$\partial q_{\text{all}, i} / \partial S_{\text{merge}} = 0$である.

需要と供給が等しい場合の微分

この場合,$q_{\text{all}, i} = D_i$である.
ここは不連続点であるため,微分の値を適当に決める必要がある.

ただし以下の条件を満たすようにする:
$$q_{\text{out}} = \sum_{i=1}^N q_{\text{all}, i}$$
$$\frac{\partial q_{\text{out}}}{\partial S_{\text{merge}}} = \sum_{i=1}^N \frac{\partial q_{\text{all}, i}}{\partial S_{\text{merge}}}$$

このときの値は不定であるため,他の実装に合わせてコメントで注記しておいたほうがよい.

供給が0の場合の微分

この場合常に渋滞であるため,$q_{\text{all}, i} = q_{\text{cong}, i}$である.
このときの微分は以下の通りである:
$$\frac{\partial q_{\text{all}, i}}{\partial S_{\text{merge}}} = \alpha_i$$
$$\frac{\partial q_{\text{all}, i}}{\partial D_j} = 0$$

これは割り算が発生しないため,供給が0の場合でも上流リンクの勾配が消えないことを意味する.

完全に渋滞している場合の微分

これは総需要が供給を超える場合,常に渋滞しているため,$q_{\text{all}, i} = q_{\text{cong}, i}$である.
このときの微分は以下の通りである:
$$\frac{\partial q_{\text{all}, i}}{\partial D_j} = 0$$
また,総流量に対する供給の微分は
$$\frac{\partial q_{\text{out}}}{\partial S_{\text{merge}}} = 1$$
である.

@toruseo

toruseo commented Sep 28, 2026

Copy link
Copy Markdown
Owner

せっかく更新していただいたのに確認遅れててすいません.こちら側で色々立て込んでまして・・・
いずれ必ず確認しますので,しばしお待ちください.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

JAX node model: merge nodes with 3+ inlinks never discharge the 3rd and later inlinks

2 participants