Repository navigation
Add correctness tests for Euler/RK4 integration - #121
anar-rzayev wants to merge 2 commits into
Conversation
|
Warning Review limit reachedNext included review available in 38 minutes. View limit detailsLimit details: You’ve used the included review currently available. You've used all free OSS reviews for now. Wait for the free limit to reset to keep reviewing this public repository. Review configuration: ⚙️ Run configurationConfiguration used: defaults Review profile: CHILL Plan: Team Run ID: 📒 Files selected for processing (1)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
There was a problem hiding this comment.
Pull request overview
Adds mathematically grounded integration correctness tests to ensure FlowMatcher’s Euler and RK4 steppers produce expected trajectories under analytic velocity fields, addressing Issue #58 and strengthening confidence beyond shape/key assertions.
Changes:
- Introduces small analytic velocity-field
torch.nn.Modulehelpers for constant, linear decay, time-ramp, and per-graph velocities. - Adds correctness tests for Euler/RK4 against closed-form expectations (constant field, time-dependent stage timing) and relative convergence (RK4 vs Euler).
- Adds a batched-graphs test to verify per-graph result/trajectory splitting is correct when water counts differ.
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| def run(method: str, num_steps: int, seed: int): | ||
| # Same seed => same prior noise (the integration's only randomness), | ||
| # so the runs are compared from identical initial positions. | ||
| gen = torch.Generator(device=device).manual_seed(seed) | ||
| integrate = getattr(flow_matcher, f"{method}_integrate") | ||
| return integrate( | ||
| simple_hetero_data, | ||
| num_steps=num_steps, | ||
| device=str(device), | ||
| return_trajectory=True, | ||
| generator=gen, | ||
| )[0] | ||
|
|
||
| seed = 1234 | ||
| euler_fine = run("euler", num_steps=2000, seed=seed) | ||
| rk4_coarse = run("rk4", num_steps=21, seed=seed) | ||
| euler_coarse = run("euler", num_steps=21, seed=seed) | ||
|
|
||
| # Identical initial noise across the three runs (guards the comparison). | ||
| np.testing.assert_allclose( | ||
| rk4_coarse["trajectory"][0], euler_fine["trajectory"][0], atol=1e-5 | ||
| ) | ||
| np.testing.assert_allclose( | ||
| euler_coarse["trajectory"][0], euler_fine["trajectory"][0], atol=1e-5 | ||
| ) |
Another go at #58.
test_euler_integrateandtest_rk4_integrateassert shapes and dict keys only, so nothing in the suite notices if either integrator computes the wrong numbers.The approach is to swap the model for an analytic velocity field. The ODE then has a solution in closed form, and the integrator output can be compared to it directly - this tests the stepper arithmetic, not the network.
Constant field,
v = c(both methods). Integratingdx/dt = covert ∈ [0,1]givesx(1) = x(0) + cwhatever the step count,and frame
ksits onx0 + k/(N-1) * c. RK4 is exact here too, since(dt/6)(c + 2c + 2c + c) = dt*c.Decay field,
v = -x(the convergence check). The solution isx(1) = x(0) * e^-1. Against a 2000-step Euler reference, 21-step RK4 agrees tortol=2e-2, and its error is strictly smaller than 21-step Euler's - the fourthorder buys accuracy the first order does not.
Ramp field,
v = a*t. Both fields above ignoret, so neither notices astage evaluated at the wrong time - but the model is time-conditioned and RK4's
midpoint and endpoint times are load-bearing. Here the quadrature is known:
Simpson's weights are exact on a linear integrand, so RK4 lands on
a/2, whileEuler's left rectangles give exactly
a*(N-2)/(2*(N-1)).Batching. Two graphs with 5 and 3 waters and a per-graph velocity
(i+1)*c, covering the list-of-graphs path inscripts/inference.pythat everyexisting integration test misses.