Repository navigation
Conversation
…3rd and later inlinks Fixes toruseo#18 and add test
|
@Antikira Issueのほうに書いてある通り,この修正にはdesign decisionが必要です. あと,本リポジトリのライセンスをMITからApache 2.0に変更します.Antikira さんの貢献もApacheになることになりますが,よろしいでしょうか? |
|
@toruseo 今回の実装方針を採用した理由は主に以下の2点です.
割り当て規則の同等性を保ちつつ,速度面でもメリットが見込めるため,今回はこちらの実装方式を選択しています. ライセンス変更の件は同意いたします. |
|
ありがとうございます.自分もその設計方針に賛成です. astra先生にレビューしてもらったら以下のコメントをもらいました.妥当に思えるので,修正お願いします. 対象はコミット 1. 【P1】優先配分量と需要が一致する点で、総流量の勾配が誤る箇所: このPRは、既存の2本合流の中央値式も、3本以上に対応する式へ置き換えています。置き換え後の式で、流量は一致しても勾配が誤るケースがあります。(GitHub) 例えば、2本の流入リンクについて、需要を 総流量を (Q=q_1+q_2) とすると、この近傍では供給制約が有効で、(Q=S) が成立します。該当計算式を単体実行した結果は以下です。
float32・float64、JITあり・なしで同じ結果を再現しました。3本合流でも、 この例では、 修正要求: 2本合流は既存の中央値式を保持するのが最小変更です。3本以上の合流については、Python側の配分を保ちつつ、供給制約下で総流量の微分が「需要について0、優先度について0、供給について1」になるよう計算式を修正する必要があります。 2. 【P2】正規化分母の下限により、Python側と配分が変わる箇所: Python側は優先度の総和で正規化しています。このPRでは総和に 例えば、 優先度を同じ比率の 修正要求: 正の総優先度には、その値を分母として使う必要があります。パディングによって総優先度がゼロになる行は、分母を別途安全化できます。 3. 追加テストは、リンク別の挙動と勾配まで検証する必要がある追加された 総旅行時間という集計値だけでは、リンク間の配分の違いを捉えきれません。また、この追加テストには自動微分の検査が含まれていません。 ご指定の設計判断に対応するテストとして、次の3点が必要です。
設計方針についての確認結果Python側の「優先度による初期配分と、登録順での余剰供給の再配分」をJAX側へ移植する方針は、挙動統一という目的に対応しています。 今回、優先度を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)。リポジトリ全体の |
…ts for merge nodes
|
修正しました. 修正内容P1関連: 3cdfabc最大化と最小化関数が入力値が等しい場合に微分を0.5に割り当てる仕様のため発生していました. 詳細な検討はAppendixのとおりです. P2関連: b845183もとの実装は0除算エラー対策です. P3 関連: a39da64, a46d993, cc7af01, 3390d9e
Appendixjaxの勾配に関して以下のコードで勾配を追跡したところ局地処理が問題の原因だと特定しました. # 誤差の検証
# 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}")証明となる. 渋滞中は となる.この総和が個別の流量の和が総流量であることを念頭に置くと線形性から以下が成立する. ここで, |
|
ありがとうございます. テストコードが乱立していて旧記述なども紛れていているので,いったん整理していただけますか? あと,例によってastraが以下の問題を見つけてます:
追記:全体的方針はよいと思います |
There was a problem hiding this comment.
修正しました.
テストコードの肥大化に関しては間違えて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}$は
となる.
このときの未配分余剰供給
である.
各リンクの未充足需要$\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}$が満たすべき条件は以下の通りである:
- 流量がもとの値と変わらないこと:$q_{\text{cong}, i} = q_{\text{raw}, i}$
- 各リンクの需要が供給の優先配分を超える場合において,流量の供給微分$\partial q_{\text{cong}, i} / \partial S_{\text{merge}}=\alpha_i$かつ需要微分$\partial q_{\text{cong}, i} / \partial D_j=0$が成立すること.
- 供給が0の場合においては供給による微分が1であり,かつ需要による微分が0であること.
色々探したところ以下の式が良さそうである:
渋滞時において,$\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{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$である.
ここは不連続点であるため,微分の値を適当に決める必要がある.
ただし以下の条件を満たすようにする:
このときの値は不定であるため,他の実装に合わせてコメントで注記しておいたほうがよい.
供給が0の場合の微分
この場合常に渋滞であるため,$q_{\text{all}, i} = q_{\text{cong}, i}$である.
このときの微分は以下の通りである:
これは割り算が発生しないため,供給が0の場合でも上流リンクの勾配が消えないことを意味する.
完全に渋滞している場合の微分
これは総需要が供給を超える場合,常に渋滞しているため,$q_{\text{all}, i} = q_{\text{cong}, i}$である.
このときの微分は以下の通りである:
また,総流量に対する供給の微分は
である.
|
せっかく更新していただいたのに確認遅れててすいません.こちら側で色々立て込んでまして・・・ |
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.pyVectorized the multi-inlink merge priority logic from
Node._transfer_merge(unsim/unsim.py) into a vectorize:2.
tests/test_diff.pytest_merge_3inlinksto verify that Python and JAX simulation total travel time match exactly on a 3-inlink congested merge network.