Initial commit

This commit is contained in:
rossiyareich
2024-07-21 23:00:34 +07:00
parent 0df74170a1
commit be0e6fd0bb
31 changed files with 1898 additions and 0 deletions
+2
View File
@@ -160,3 +160,5 @@ cython_debug/
# and can be added to the global gitignore or merged into this file. For a more nuclear
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
#.idea/
.vscode/
+16
View File
@@ -0,0 +1,16 @@
"""
@author: wootwootwootwoot
@title: ComfyUI-RK-Sampler
@nickname: ComfyUI-RK-Sampler
@description: Batched Runge-Kutta Samplers for ComfyUI
"""
from .nodes import nodes_rk_sampler
NODE_CLASS_MAPPINGS = {
**nodes_rk_sampler.NODE_CLASS_MAPPINGS,
}
NODE_DISPLAY_NAME_MAPPINGS = {
**nodes_rk_sampler.NODE_DISPLAY_NAME_MAPPINGS,
}
+3
View File
@@ -0,0 +1,3 @@
{
"RungeKuttaSampler": ""
}
View File
View File
+29
View File
@@ -0,0 +1,29 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class AEBosh3(ExplicitRungeKutta):
ORDER = 3
NFE_PER_STEP = 3
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 2, 3 / 4, 1],
a=[[], [1 / 2], [0, 3 / 4], [2 / 9, 1 / 3, 4 / 9]],
b=[2 / 9, 1 / 3, 4 / 9, 0],
b_err=[2 / 9 - 7 / 24, 1 / 3 - 1 / 4, 4 / 9 - 1 / 3, 0 - 1 / 8],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+43
View File
@@ -0,0 +1,43 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class AECashKarp5(ExplicitRungeKutta):
ORDER = 5
NFE_PER_STEP = 6
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 5, 3 / 10, 3 / 5, 1, 7 / 8],
a=[
[],
[1 / 5],
[3 / 40, 9 / 40],
[3 / 10, -9 / 10, 6 / 5],
[-11 / 54, 5 / 2, -70 / 27, 35 / 27],
[1631 / 55296, 175 / 512, 575 / 13824, 44275 / 110592, 253 / 4096],
],
b=[37 / 378, 0, 250 / 621, 125 / 594, 0, 512 / 1771],
b_err=[
37 / 378 - 2825 / 27648,
0 - 0,
250 / 621 - 18575 / 48384,
125 / 594 - 13525 / 55296,
0 - 277 / 14336,
512 / 1771 - 1 / 4,
],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+58
View File
@@ -0,0 +1,58 @@
from typing import Optional
import torch
from torchode.interpolation import FourthOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class AEDopri5(ExplicitRungeKutta):
ORDER = 5
NFE_PER_STEP = 6
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 5, 3 / 10, 4 / 5, 8 / 9, 1, 1],
a=[
[],
[1 / 5],
[3 / 40, 9 / 40],
[44 / 45, -56 / 15, 32 / 9],
[19372 / 6561, -25360 / 2187, 64448 / 6561, -212 / 729],
[9017 / 3168, -355 / 33, 46732 / 5247, 49 / 176, -5103 / 18656],
[35 / 384, 0, 500 / 1113, 125 / 192, -2187 / 6784, 11 / 84],
],
b=[35 / 384, 0, 500 / 1113, 125 / 192, -2187 / 6784, 11 / 84, 0],
b_err=[
35 / 384 - 5179 / 57600,
0 - 0,
500 / 1113 - 7571 / 16695,
125 / 192 - 393 / 640,
-2187 / 6784 - -92097 / 339200,
11 / 84 - 187 / 2100,
0 - 1 / 40,
],
b_other=[
[
6025192743 / 30085553152 / 2,
0,
51252292925 / 65400821598 / 2,
-2691868925 / 45128329728 / 2,
187940372067 / 1594534317056 / 2,
-1776094331 / 19743644256 / 2,
11237099 / 235043384 / 2,
]
],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return FourthOrderPolynomialInterpolation.from_k(
data.t0, data.dt, data.y0, data.y1, data.k, data.tableau.b_other[0]
)
+264
View File
@@ -0,0 +1,264 @@
from typing import Optional
import torch
from torchode.interpolation import FourthOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class AEDopri8(ExplicitRungeKutta):
ORDER = 8
NFE_PER_STEP = 13
TABLEAU = ButcherTableau.from_lists(
c=[
0,
1 / 18,
1 / 12,
1 / 8,
5 / 16,
3 / 8,
59 / 400,
93 / 200,
5490023248 / 9719169821,
13 / 20,
1201146811 / 1299019798,
1,
1,
1,
],
a=[
[],
[1 / 18],
[1 / 48, 1 / 16],
[1 / 32, 0, 3 / 32],
[5 / 16, 0, -75 / 64, 75 / 64],
[3 / 80, 0, 0, 3 / 16, 3 / 20],
[
29443841 / 614563906,
0,
0,
77736538 / 692538347,
-28693883 / 1125000000,
23124283 / 1800000000,
],
[
16016141 / 946692911,
0,
0,
61564180 / 158732637,
22789713 / 633445777,
545815736 / 2771057229,
-180193667 / 1043307555,
],
[
39632708 / 573591083,
0,
0,
-433636366 / 683701615,
-421739975 / 2616292301,
100302831 / 723423059,
790204164 / 839813087,
800635310 / 3783071287,
],
[
246121993 / 1340847787,
0,
0,
-37695042795 / 15268766246,
-309121744 / 1061227803,
-12992083 / 490766935,
6005943493 / 2108947869,
393006217 / 1396673457,
123872331 / 1001029789,
],
[
-1028468189 / 846180014,
0,
0,
8478235783 / 508512852,
1311729495 / 1432422823,
-10304129995 / 1701304382,
-48777925059 / 3047939560,
15336726248 / 1032824649,
-45442868181 / 3398467696,
3065993473 / 597172653,
],
[
185892177 / 718116043,
0,
0,
-3185094517 / 667107341,
-477755414 / 1098053517,
-703635378 / 230739211,
5731566787 / 1027545527,
5232866602 / 850066563,
-4093664535 / 808688257,
3962137247 / 1805957418,
65686358 / 487910083,
],
[
403863854 / 491063109,
0,
0,
-5068492393 / 434740067,
-411421997 / 543043805,
652783627 / 914296604,
11173962825 / 925320556,
-13158990841 / 6184727034,
3936647629 / 1978049680,
-160528059 / 685178525,
248638103 / 1413531060,
0,
],
[
14005451 / 335480064,
0,
0,
0,
0,
-59238493 / 1068277825,
181606767 / 758867731,
561292985 / 797845732,
-1041891430 / 1371343529,
760417239 / 1151165299,
118820643 / 751138087,
-528747749 / 2220607170,
1 / 4,
],
],
b=[
14005451 / 335480064,
0,
0,
0,
0,
-59238493 / 1068277825,
181606767 / 758867731,
561292985 / 797845732,
-1041891430 / 1371343529,
760417239 / 1151165299,
118820643 / 751138087,
-528747749 / 2220607170,
1 / 4,
0,
],
b_err=[
14005451 / 335480064 - 13451932 / 455176623,
0 - 0,
0 - 0,
0 - 0,
0 - 0,
-59238493 / 1068277825 - -808719846 / 976000145,
181606767 / 758867731 - 1757004468 / 5645159321,
561292985 / 797845732 - 656045339 / 265891186,
-1041891430 / 1371343529 - -3867574721 / 1518517206,
760417239 / 1151165299 - 465885868 / 322736535,
118820643 / 751138087 - 53011238 / 667516719,
-528747749 / 2220607170 - 2 / 45,
1 / 4 - 0,
0 - 0,
],
b_other=[
[
(
-6.3448349392860401388 * (0.5**5)
+ 22.1396504998094068976 * (0.5**4)
- 30.0610568289666450593 * (0.5**3)
+ 19.9990069333683970610 * (0.5**2)
- 6.6910181737837595697 * 0.5
+ 1.0
)
/ 2,
0,
0,
0,
0,
(
-39.6107919852202505218 * (0.5**5)
+ 116.4422149550342161651 * (0.5**4)
- 121.4999627731334642623 * (0.5**3)
+ 52.2273532792945524050 * (0.5**2)
- 7.6142658045872677172 * 0.5
)
/ 2,
(
20.3761213808791436958 * (0.5**5)
- 67.1451318825957197185 * (0.5**4)
+ 83.1721004639847717481 * (0.5**3)
- 46.8919164181093621583 * (0.5**2)
+ 10.7281392630428866124 * 0.5
)
/ 2,
(
7.3347098826795362023 * (0.5**5)
- 16.5672243527496524646 * (0.5**4)
+ 9.5724507555993664382 * (0.5**3)
- 0.1890893225010595467 * (0.5**2)
+ 0.5526637063753648783 * 0.5
)
/ 2,
(
32.8801774352459155182 * (0.5**5)
- 89.9916014847245016028 * (0.5**4)
+ 87.8406057677205645007 * (0.5**3)
- 35.7075975946222072821 * (0.5**2)
+ 4.2186562625665153803 * 0.5
)
/ 2,
(
-10.1588990526426760954 * (0.5**5)
+ 22.6237489648532849093 * (0.5**4)
- 17.4152107770762969005 * (0.5**3)
+ 6.2736448083240352160 * (0.5**2)
- 0.6627209125361597559 * 0.5
)
/ 2,
(
-12.5401268098782561200 * (0.5**5)
+ 32.2362340167355370113 * (0.5**4)
- 28.5903289514790976966 * (0.5**3)
+ 10.3160881272450748458 * (0.5**2)
- 1.2636789001135462218 * 0.5
)
/ 2,
(
29.5553001484516038033 * (0.5**5)
- 82.1020315488359848644 * (0.5**4)
+ 81.6630950584341412934 * (0.5**3)
- 34.7650769866611817349 * (0.5**2)
+ 5.4106037898590422230 * 0.5
)
/ 2,
(
-41.7923486424390588923 * (0.5**5)
+ 116.2662185791119533462 * (0.5**4)
- 114.9375291377009418170 * (0.5**3)
+ 47.7457971078225540396 * (0.5**2)
- 7.0321379067945741781 * 0.5
)
/ 2,
(
20.3006925822100825485 * (0.5**5)
- 53.9020777466385396792 * (0.5**4)
+ 50.2558364226176017553 * (0.5**3)
- 19.0082099341608028453 * (0.5**2)
+ 2.3537586759714983486 * 0.5
)
/ 2,
]
],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return FourthOrderPolynomialInterpolation.from_k(
data.t0, data.dt, data.y0, data.y1, data.k, data.tableau.b_other[0]
)
+29
View File
@@ -0,0 +1,29 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class AEFehlberg2(ExplicitRungeKutta):
ORDER = 2
NFE_PER_STEP = 3
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 2, 1],
a=[[], [1 / 2], [1 / 256, 255 / 256]],
b=[1 / 512, 255 / 256, 1 / 512],
b_err=[1 / 512 - 1 / 256, 255 / 256 - 255 / 256, 1 / 512 - 0],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+43
View File
@@ -0,0 +1,43 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class AEFehlberg5(ExplicitRungeKutta):
ORDER = 5
NFE_PER_STEP = 6
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 4, 3 / 8, 12 / 13, 1, 1 / 2],
a=[
[],
[1 / 4],
[3 / 32, 9 / 32],
[1932 / 2197, -7200 / 2197, 7296 / 2197],
[439 / 216, -8, 3680 / 513, -845 / 4104],
[-8 / 27, 2, -3544 / 2565, 1859 / 4104, -11 / 40],
],
b=[16 / 135, 0, 6656 / 12825, 28561 / 56430, -9 / 50, 2 / 55],
b_err=[
16 / 135 - 25 / 216,
0 - 0,
6656 / 12825 - 1408 / 2565,
28561 / 56430 - 2197 / 4104,
-9 / 50 - -1 / 5,
2 / 55 - 0,
],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+29
View File
@@ -0,0 +1,29 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class AEHeunEuler2(ExplicitRungeKutta):
ORDER = 2
NFE_PER_STEP = 2
TABLEAU = ButcherTableau.from_lists(
c=[0, 1],
a=[[], [1]],
b=[1 / 2, 1 / 2],
b_err=[1 / 2 - 1, 1 / 2 - 0],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+29
View File
@@ -0,0 +1,29 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class AEMidpoint2(ExplicitRungeKutta):
ORDER = 2
NFE_PER_STEP = 2
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 2],
a=[[], [1 / 2]],
b=[0, 1],
b_err=[0 - -1, 1 - 2],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+29
View File
@@ -0,0 +1,29 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class AERalston2(ExplicitRungeKutta):
ORDER = 2
NFE_PER_STEP = 2
TABLEAU = ButcherTableau.from_lists(
c=[0, 2 / 3],
a=[[], [2 / 3]],
b=[1 / 4, 3 / 4],
b_err=[1 / 4 - -2 / 4, 3 / 4 - 6 / 4],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+134
View File
@@ -0,0 +1,134 @@
from typing import Optional
import sympy as sp
import torch
from torchode.interpolation import FourthOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
def compute_interpolation_weights():
"""Compute the interpolation weights for the Tsit5 interpolation coefficients of
2nd, 3rd and 4th order.
The original Tsit5 paper builds the 4th order interpolation polynomial as a linear
combination of 7 polynomials. This function computes weights that give us the
coefficients a, b, c in the standard polynomial form `a*x^4 + b*x^3 + ...` directly.
This way, we can evaluate the interpolant more efficiently.
"""
t, y0, dt, f1, f2, f3, f4, f5, f6, f7 = sp.symbols("t y_0 dt f_1 f_2 f_3 f_4 f_5 f_6 f_7")
f = [f1, f2, f3, f4, f5, f6, f7]
# fmt: off
# The 7 basis functions of the interpolant
b = [
-1.0530884977290216 * t * (t - 1.3299890189751412) * (t**2 - 1.4364028541716351 * t + 0.7139816917074209),
0.1017 * t**2 * (t**2 - 2.1966568338249754 * t + 1.2949852507374631),
2.490627285651252793 * t**2 * (t**2 - 2.38535645472061657 * t + 1.57803468208092486),
-16.54810288924490272 * (t - 1.21712927295533244) * (t - 0.61620406037800089) * t**2,
47.37952196281928122 * (t - 1.203071208372362603) * (t - 0.658047292653547382) * t**2,
-34.87065786149660974 * (t - 1.2) * (t - 0.666666666666666667) * t**2,
2.5 * (t - 1) * (t - 0.6) * t**2
]
# fmt: on
interpolant = y0 + dt * sum(f_i * b_i for f_i, b_i in zip(f, b))
# Fully expand the polynomial and collect the powers of t
form = sp.collect(sp.expand(interpolant, t), t)
# The coefficients of t^2, t^3 and t^4 are of the form
#
# dt * \sum_i x_i f_i
#
# and we collect the x_i here in a matrix. Then the coefficients of the interpolant
# can be found by a matrix multiplication between this and the k vector of the RK
# stages.
return [[float(form.coeff(t, i).coeff(f[j]).coeff(dt)) for j in range(len(f))] for i in range(2, 5)]
class AETsit5(ExplicitRungeKutta):
"""The 5th order Runge-Kutta method by Tsitouras
References
----------
```bibtex
@article{tsitouras2011runge,
title={Runge--Kutta pairs of order 5(4) satisfying only the first column
simplifying assumption},
author={Tsitouras, Charalampos},
journal={Computers \\& Mathematics with Applications},
volume={62},
number={2},
pages={770--775},
year={2011},
publisher={Elsevier}
}
```
"""
ORDER = 5
NFE_PER_STEP = 6
TABLEAU = ButcherTableau.from_lists(
c=[0.0, 0.161, 0.327, 0.9, 0.9800255409045097, 1.0, 1.0],
a=[
# fmt: off
[],
[0.161],
[-0.008480655492356989, 0.335480655492357],
[2.8971530571054935, -6.359448489975075, 4.3622954328695815],
[5.325864828439257, -11.748883564062828, 7.4955393428898365, -0.09249506636175525],
[5.86145544294642, -12.92096931784711, 8.159367898576159, -0.071584973281401, -0.02826905039406838],
[0.09646076681806523, 0.01, 0.4798896504144996, 1.379008574103742, -3.290069515436081, 2.324710524099774],
# fmt: on
],
b=[
0.09646076681806523,
0.01,
0.4798896504144996,
1.379008574103742,
-3.290069515436081,
2.324710524099774,
0.0,
],
# The paper introduces b-tilde as the weights of the lower-order interpolant but
# the weights they give in the end are actually directly the weights for the
# error estimate, see [1].
#
# [1] https://github.com/patrick-kidger/diffrax/issues/98
b_err=[
0.00178001105222577714,
0.0008164344596567469,
-0.007880878010261995,
0.1447110071732629,
-0.5823571654525552,
0.45808210592918697,
# The original paper has the sign of this coefficient wrong, see [1]
#
# [1] https://github.com/SciML/OrdinaryDiffEq.jl/issues/1654
-1 / 66,
],
b_other=compute_interpolation_weights(),
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
y0 = data.y0
dt = data.dt.to(dtype=y0.dtype)
f0 = data.k[0]
b = data.tableau.b_other
assert b is not None
B = torch.einsum("b, cs, sbf -> cbf", dt, b, data.k)
c, b, a = B[0], B[1], B[2]
d = dt[:, None] * f0
e = y0
coefficients = (e, d, c, b, a)
return FourthOrderPolynomialInterpolation(data.t0, data.t0 + data.dt, coefficients)
+28
View File
@@ -0,0 +1,28 @@
from typing import Optional
import torch
from torchode.interpolation import LinearInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class FEEuler1(ExplicitRungeKutta):
ORDER = 1
NFE_PER_STEP = 1
TABLEAU = ButcherTableau.from_lists(
c=[0],
a=[[]],
b=[1],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return LinearInterpolation(data.t0, data.dt, data.y0, data.y1)
+28
View File
@@ -0,0 +1,28 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class FEHeun3(ExplicitRungeKutta):
ORDER = 3
NFE_PER_STEP = 3
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 3, 2 / 3],
a=[[], [1 / 3], [0, 2 / 3]],
b=[1 / 4, 0, 3 / 4],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+28
View File
@@ -0,0 +1,28 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class FEKutta3(ExplicitRungeKutta):
ORDER = 3
NFE_PER_STEP = 3
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 2, 1],
a=[[], [1 / 2], [-1, 2]],
b=[1 / 6, 2 / 3, 1 / 6],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+28
View File
@@ -0,0 +1,28 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class FEKutta4(ExplicitRungeKutta):
ORDER = 4
NFE_PER_STEP = 4
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 2, 1 / 2, 1],
a=[[], [1 / 2], [0, 1 / 2], [0, 0, 1]],
b=[1 / 6, 1 / 3, 1 / 3, 1 / 6],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+28
View File
@@ -0,0 +1,28 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class FEKutta38th4(ExplicitRungeKutta):
ORDER = 4
NFE_PER_STEP = 4
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 3, 2 / 3, 1],
a=[[], [1 / 3], [-1 / 3, 1], [1, -1, 1]],
b=[1 / 8, 3 / 8, 3 / 8, 1 / 8],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+28
View File
@@ -0,0 +1,28 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class FERalston3(ExplicitRungeKutta):
ORDER = 3
NFE_PER_STEP = 3
TABLEAU = ButcherTableau.from_lists(
c=[0, 1 / 2, 3 / 4],
a=[[], [1 / 2], [0, 3 / 4]],
b=[2 / 9, 1 / 3, 4 / 9],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+41
View File
@@ -0,0 +1,41 @@
from math import sqrt
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ExplicitRungeKutta, ERKInterpolationData
SQRT5 = sqrt(5)
class FERalston4(ExplicitRungeKutta):
ORDER = 4
NFE_PER_STEP = 4
TABLEAU = ButcherTableau.from_lists(
c=[0, 2 / 5, (14 - 3 * SQRT5) / 16, 1],
a=[
[],
[2 / 5],
[(-2889 + 1428 * SQRT5) / 1024, (3785 - 1620 * SQRT5) / 1024],
[(-3365 + 2094 * SQRT5) / 6040, (-975 - 3046 * SQRT5) / 2552, (467040 + 203968 * SQRT5) / 240845],
],
b=[
(263 + 24 * SQRT5) / 1812,
(125 - 1000 * SQRT5) / 3828,
(3426304 + 1661952 * SQRT5) / 5924787,
(30 - 4 * SQRT5) / 123,
],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+28
View File
@@ -0,0 +1,28 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class FESSPRK3(ExplicitRungeKutta):
ORDER = 3
NFE_PER_STEP = 3
TABLEAU = ButcherTableau.from_lists(
c=[0, 1, 1 / 2],
a=[[], [1], [1 / 4, 1 / 4]],
b=[1 / 6, 1 / 6, 2 / 3],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+28
View File
@@ -0,0 +1,28 @@
from typing import Optional
import torch
from torchode.interpolation import ThirdOrderPolynomialInterpolation
from torchode.terms import ODETerm
from .runge_kutta import ButcherTableau, ERKInterpolationData, ExplicitRungeKutta
class FEWray3(ExplicitRungeKutta):
ORDER = 3
NFE_PER_STEP = 3
TABLEAU = ButcherTableau.from_lists(
c=[0, 8 / 15, 2 / 3],
a=[[], [8 / 15], [1 / 4, 5 / 12]],
b=[1 / 4, 0, 3 / 4],
)
def __init__(self, term: Optional[ODETerm] = None):
super().__init__(term, self.TABLEAU)
@torch.jit.export
def convergence_order(self):
return self.ORDER
@torch.jit.export
def build_interpolation(self, data: ERKInterpolationData):
return ThirdOrderPolynomialInterpolation.from_k(data.t0, data.dt, data.y0, data.y1, data.k)
+249
View File
@@ -0,0 +1,249 @@
from typing import Any, Dict, List, Optional, Tuple
import torch
import torch.nn as nn
from torchode.interpolation import LocalInterpolation
from torchode.problems import InitialValueProblem
from torchode.single_step_methods.base import StepResult
from torchode.single_step_methods.runge_kutta import (
CoefficientVector,
ERKInterpolationData,
ERKState,
RungeKuttaMatrix,
WeightMatrix,
WeightVector,
)
from torchode.terms import ODETerm
from torchode.typing import *
class ButcherTableau:
def __init__(
self,
# Coefficients for the evaluation nodes in time
c: CoefficientVector,
# Runge-Kutta matrix
a: RungeKuttaMatrix,
# Coefficients for the high-order solution estimate
b: WeightVector,
# Coefficients for the error estimate
b_err: WeightVector,
# Additional additional rows of the b matrix
b_other: Optional[WeightMatrix] = None,
fsal: Optional[bool] = None,
ssal: Optional[bool] = None,
):
self.c = c
self.a = a
self.b = b
self.b_err = b_err
self.b_other = b_other
if fsal is None:
fsal = self.is_fsal()
self.fsal = fsal
if ssal is None:
ssal = self.is_ssal()
self.ssal = ssal
@staticmethod
def from_lists(
*,
c: List[float],
a: List[List[float]],
b: List[float],
b_err: Optional[List[float]] = None,
b_low_order: Optional[List[float]] = None,
b_other: Optional[List[List[float]]] = None,
):
is_adaptive = b_err is not None or b_low_order is not None
n_nodes = len(c)
n_weights = len(b)
assert n_nodes == n_weights
assert len(a) == n_nodes
# Fill a up into a full square matrix
a_full = [row + [0.0] * (n_weights - len(row)) for row in a]
b_coeffs = torch.tensor(b, dtype=torch.float64)
b_err_coeffs = None
if is_adaptive:
if b_err is None:
assert b_low_order is not None
assert len(b_low_order) == n_weights
b_low_coeffs = torch.tensor(b_low_order, dtype=torch.float64)
b_err_coeffs = b_coeffs - b_low_coeffs
else:
b_err_coeffs = torch.tensor(b_err, dtype=torch.float64)
if b_other is None:
b_other_coeffs = None
else:
b_other_coeffs = torch.tensor(b_other, dtype=torch.float64)
assert b_other_coeffs.ndim == 2
assert b_other_coeffs.shape[1] == n_weights
return ButcherTableau(
c=torch.tensor(c, dtype=torch.float64),
a=torch.tensor(a_full, dtype=torch.float64),
b=b_coeffs,
b_err=b_err_coeffs,
b_other=b_other_coeffs,
)
def to(self, device: torch.device, time_dtype: torch.dtype, data_dtype: torch.dtype) -> "ButcherTableau":
b_other = self.b_other
if b_other is not None:
b_other = b_other.to(device, data_dtype)
b_err = None
if self.b_err is not None:
b_err = self.b_err.to(device, data_dtype)
return ButcherTableau(
c=self.c.to(device, time_dtype),
a=self.a.to(device, data_dtype),
b=self.b.to(device, data_dtype),
b_err=b_err,
b_other=b_other,
fsal=self.fsal,
ssal=self.ssal,
)
@property
def n_stages(self):
return self.c.shape[0]
def is_fsal(self):
"""Is `f(y0)` equal to `f(y1)` from the previous step?
If that is the case, we can reuse the result from the previous step.
"""
is_lower_triangular = (torch.triu(self.a, diagonal=1) == 0.0).all().item()
first_node_is_t0 = (self.c[0] == 0.0).item()
last_node_is_t1 = (self.c[-1] == 1.0).item()
first_stage_explicit = (self.a[0, 0] == 0.0).item()
return (
is_lower_triangular
and (self.b == self.a[-1]).all().item()
and first_node_is_t0
and last_node_is_t1
and first_stage_explicit
)
def is_ssal(self):
"""Is the solution equal to the last stage result?
If that is the case, we can avoid the final computation of the solution and
return the last stage result instead.
"""
is_lower_triangular = (torch.triu(self.a, diagonal=1) == 0.0).all().item()
last_node_is_t1 = (self.c[-1] == 1.0).item()
last_stage_explicit = (self.a[-1, -1] == 0.0).item()
return is_lower_triangular and (self.b == self.a[-1]).all().item() and last_node_is_t1 and last_stage_explicit
class ExplicitRungeKutta(nn.Module):
def __init__(self, term: Optional[ODETerm], tableau: ButcherTableau):
super().__init__()
self.term = term
self.tableau = tableau
@torch.jit.export
def init(
self,
term: Optional[ODETerm],
problem: InitialValueProblem,
f0: Optional[DataTensor],
*,
stats: Dict[str, Any],
args: Any,
) -> ERKState:
if self.tableau.fsal:
term_ = term
if torch.jit.is_scripting() or term_ is None:
assert term is None, "The integration term is fixed for JIT compilation"
term_ = self.term
assert term_ is not None
if f0 is None:
prev_vf1 = term_.vf(problem.t_start, problem.y0, stats, args)
else:
prev_vf1 = f0
else:
prev_vf1 = None
return ERKState(
tableau=self.tableau.to(
device=problem.device,
data_dtype=problem.data_dtype,
time_dtype=problem.time_dtype,
),
prev_vf1=prev_vf1,
)
@torch.jit.export
def merge_states(self, accept: AcceptTensor, current: ERKState, previous: ERKState):
prev_vf1 = previous.prev_vf1
current_vf1 = current.prev_vf1
if current_vf1 is None or prev_vf1 is None:
return current
else:
return ERKState(current.tableau, torch.where(accept[:, None], current_vf1, prev_vf1))
@torch.jit.export
def step(
self,
term: Optional[ODETerm],
running: AcceptTensor,
y0: DataTensor,
t0: TimeTensor,
dt: TimeTensor,
state: ERKState,
*,
stats: Dict[str, Any],
args: Any,
) -> Tuple[StepResult, ERKInterpolationData, ERKState, Optional[StatusTensor]]:
term_ = term
if torch.jit.is_scripting() or term_ is None:
assert term is None, "The integration term is fixed for JIT compilation"
term_ = self.term
assert term_ is not None
tableau = state.tableau
# Convert dt into the data dtype for dtype stability
dt_data = dt.to(dtype=y0.dtype)
prev_vf1 = state.prev_vf1
vf0 = prev_vf1 if tableau.fsal and prev_vf1 is not None else term_.vf(t0, y0, stats, args)
y_i = y0
k = vf0.new_empty((tableau.n_stages, vf0.shape[0], vf0.shape[1]))
k[0] = vf0
a = tableau.a
t_nodes = torch.addcmul(t0, tableau.c[:, None], dt)
for i in range(1, tableau.n_stages):
y_i = torch.einsum("j, jbf -> bf", a[i, :i], k[:i])
y_i = torch.addcmul(y0, dt_data[:, None], y_i)
k[i] = term_.vf(t_nodes[i], y_i, stats, args)
if tableau.ssal:
y1 = y_i
else:
y1 = y0 + torch.einsum("b, s, sbf -> bf", dt_data, tableau.b, k)
error_estimate = None
if tableau.b_err is not None:
error_estimate = torch.einsum("b, s, sbf -> bf", dt_data, tableau.b_err, k)
if tableau.fsal:
state = ERKState(state.tableau, k[-1])
return (
StepResult(y1, error_estimate),
ERKInterpolationData(tableau, t0, dt, y0, y1, k),
state,
None,
)
def build_interpolation(self, data: ERKInterpolationData) -> LocalInterpolation:
raise NotImplementedError()
+321
View File
@@ -0,0 +1,321 @@
import logging
import torch
import torchode
from tqdm.auto import tqdm
import comfy
import comfy.model_patcher
import comfy.samplers
import comfy.utils
from .methods.ae_bosh3 import AEBosh3
from .methods.ae_cash_karp5 import AECashKarp5
from .methods.ae_dopri5 import AEDopri5
from .methods.ae_dopri8 import AEDopri8
from .methods.ae_fehlberg2 import AEFehlberg2
from .methods.ae_fehlberg5 import AEFehlberg5
from .methods.ae_heun_euler2 import AEHeunEuler2
from .methods.ae_midpoint2 import AEMidpoint2
from .methods.ae_ralston2 import AERalston2
from .methods.ae_tsit5 import AETsit5
from .methods.fe_euler1 import FEEuler1
from .methods.fe_heun3 import FEHeun3
from .methods.fe_kutta3 import FEKutta3
from .methods.fe_kutta4 import FEKutta4
from .methods.fe_kutta_38th4 import FEKutta38th4
from .methods.fe_ralston3 import FERalston3
from .methods.fe_ralston4 import FERalston4
from .methods.fe_ssprk3 import FESSPRK3
from .methods.fe_wray3 import FEWray3
from .step_size_controllers.pid_controller import PIDController
from .step_size_controllers.scheduled_controller import ScheduledController
logger = logging.getLogger(__name__)
logger.propagate = False
logger.setLevel(logging.INFO)
sh = logging.StreamHandler()
sh.setFormatter(logging.Formatter(f"[ComfyUI-RK-Sampler] %(levelname)s - %(message)s"))
logger.addHandler(sh)
ADAPTIVE_METHODS = {
"ae_bosh3": AEBosh3,
"ae_cash_karp5": AECashKarp5,
"ae_dopri5": AEDopri5,
"ae_dopri8": AEDopri8,
"ae_fehlberg2": AEFehlberg2,
"ae_fehlberg5": AEFehlberg5,
"ae_heun_euler2": AEHeunEuler2,
"ae_midpoint2": AEMidpoint2,
"ae_ralston2": AERalston2,
"ae_tsit5": AETsit5,
}
FIXED_METHODS = {
"fe_euler1": FEEuler1,
"fe_heun3": FEHeun3,
"fe_kutta_38th4": FEKutta38th4,
"fe_kutta3": FEKutta3,
"fe_kutta4": FEKutta4,
"fe_ralston3": FERalston3,
"fe_ralston4": FERalston4,
"fe_ssprk3": FESSPRK3,
"fe_wray3": FEWray3,
}
METHODS = {**ADAPTIVE_METHODS, **FIXED_METHODS}
STEP_SIZE_CONTROLLERS = dict.fromkeys(
[
"adaptive_pid",
"fixed_scheduled",
]
)
NORMS = {
"rms_norm": torchode.step_size_controllers.rms_norm,
"max_norm": torchode.step_size_controllers.max_norm,
}
class ODETerm:
def __init__(
self,
model,
x_dtype,
x_shape,
t_dtype,
min_sigma,
t_max,
t_min,
n_steps,
is_adaptive,
method,
extra_args=None,
callback=None,
):
self.model = model
self.x_dtype = x_dtype
self.x_shape = x_shape
self.t_dtype = t_dtype
self.min_sigma = min_sigma
self.t_max = t_max
self.t_min = t_min
self.n_steps = n_steps
self.is_adaptive = is_adaptive
self.method = method
self.extra_args = {} if extra_args is None else extra_args
self.callback = callback
self.step = 0
self.nfe_step = 0
if is_adaptive:
self.progress_bar = tqdm(total=100, desc=f"Adaptive {method}", unit="%")
else:
self.progress_bar = tqdm(total=n_steps, desc=f"Fixed {method}", unit="step")
def _callback(self, t, y, denoised, mask):
if self.is_adaptive:
progress = ((self.t_max - t) / (self.t_max - self.t_min)).detach().mean().item()
d_progress = progress * 100
self.progress_bar.update(d_progress - self.step)
self.step = d_progress
i = round(progress * self.n_steps)
else:
self.progress_bar.update(1)
self.step += 1
i = self.step
if self.callback is not None:
samples = torch.where(
mask.view(*mask.shape, 1, 1, 1),
y,
denoised,
)
self.callback(
{
"x": y.to(self.x_dtype),
"i": i - 1,
"sigma": t.to(self.t_dtype),
"sigma_hat": t.to(self.t_dtype),
"denoised": samples.to(self.x_dtype),
}
)
def __call__(self, t, y):
mask = t <= self.min_sigma
y = y.reshape(self.x_shape)
denoised = torch.zeros_like(y)
if not mask.all():
denoised[~mask] = self.model(y[~mask], t[~mask], **self.extra_args)
d = torch.where(
mask.view(*mask.shape, 1, 1, 1),
torch.zeros_like(y),
(y - denoised) / t.view(*t.shape, 1, 1, 1),
)
self.nfe_step += 1
if self.nfe_step % METHODS[self.method].NFE_PER_STEP == 0:
self._callback(t, y, denoised, mask)
return d.flatten(start_dim=1)
class RungeKuttaSamplerImpl:
def __init__(
self,
method,
step_size_controller,
log_absolute_tolerance,
log_relative_tolerance,
pcoeff,
icoeff,
dcoeff,
norm,
enable_dt_min,
enable_dt_max,
dt_min,
dt_max,
safety,
factor_min,
factor_max,
max_steps,
min_sigma,
):
self.method = method
self.step_size_controller = step_size_controller
self.atol = 10**log_absolute_tolerance
self.rtol = 10**log_relative_tolerance
assert self.atol <= self.rtol
self.pcoeff = pcoeff
self.icoeff = icoeff
self.dcoeff = dcoeff
self.norm = norm
self.enable_dt_min = enable_dt_min
self.enable_dt_max = enable_dt_max
self.dt_min = dt_min
self.dt_max = dt_max
self.safety = safety
self.factor_min = factor_min
self.factor_max = factor_max
self.max_steps = max_steps
self.min_sigma = min_sigma
@torch.no_grad()
def __call__(self, model, x: torch.Tensor, sigmas: torch.Tensor, extra_args=None, callback=None, disable=None):
dtype = torch.float32 if torch.backends.mps.is_available() else torch.float64
t_max = sigmas.max()
t_min = sigmas.min()
n_steps = len(sigmas) - 1
is_adaptive = self.step_size_controller.startswith("adaptive")
if is_adaptive and (self.method in FIXED_METHODS):
raise ValueError("Fixed step methods must be used with fixed step size controllers")
term = torchode.ODETerm(
ODETerm(
model=model,
x_dtype=x.dtype,
x_shape=x.shape,
t_dtype=sigmas.dtype,
min_sigma=self.min_sigma if is_adaptive else 0.0,
t_max=t_max,
t_min=t_min,
n_steps=n_steps,
is_adaptive=is_adaptive,
method=self.method,
extra_args=extra_args,
callback=callback,
)
)
if self.step_size_controller == "fixed_scheduled":
step_size_controller = ScheduledController(sigmas=sigmas)
elif self.step_size_controller == "adaptive_pid":
step_size_controller = PIDController(
atol=self.atol,
rtol=self.rtol,
pcoeff=self.pcoeff,
icoeff=self.icoeff,
dcoeff=self.dcoeff,
term=term,
norm=NORMS[self.norm],
dt_min=self.dt_min if self.enable_dt_min else None,
dt_max=self.dt_max if self.enable_dt_max else None,
safety=self.safety,
factor_min=self.factor_min,
factor_max=self.factor_max,
)
step_method = METHODS[self.method](term=term)
adjoint = torchode.AutoDiffAdjoint(step_method, step_size_controller, max_steps=self.max_steps)
problem = torchode.InitialValueProblem(
y0=x.flatten(start_dim=1).to(dtype),
t_start=torch.full((x.shape[0],), t_max, dtype=dtype, device=sigmas.device),
t_end=torch.full((x.shape[0],), t_min, dtype=dtype, device=sigmas.device),
)
result = adjoint.solve(problem)
samples = result.ys[:, -1].reshape(x.shape).to(x.dtype)
success = True
for i, status in enumerate(result.status):
status = status.item()
if status != 0:
success = False
samples[i] = torch.full_like(samples[i], torch.nan)
reason = torchode.status_codes.Status(status)
logger.warning(f"Sample #{i} failed with reason: {reason}")
if success:
term.f.progress_bar.update(term.f.progress_bar.total - term.f.progress_bar.n)
if callback is not None:
callback(
{
"x": samples,
"i": n_steps - 1,
"sigma": t_min,
"sigma_hat": t_min,
"denoised": samples,
}
)
return samples
class RungeKuttaSampler:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"method": (list(METHODS.keys()), {"default": "ae_bosh3"}),
"step_size_controller": (list(STEP_SIZE_CONTROLLERS.keys()), {"default": "adaptive_pid"}),
"log_absolute_tolerance": ("FLOAT", {"default": -3.5}),
"log_relative_tolerance": ("FLOAT", {"default": -2.5}),
"pcoeff": ("FLOAT", {"min": 0, "default": 0.0}),
"icoeff": ("FLOAT", {"min": 0, "default": 1.0}),
"dcoeff": ("FLOAT", {"min": 0, "default": 0.0}),
"norm": (list(NORMS.keys()), {"default": "rms_norm"}),
"enable_dt_min": ("BOOLEAN", {"default": False}),
"enable_dt_max": ("BOOLEAN", {"default": True}),
"dt_min": ("FLOAT", {"default": -1.0}),
"dt_max": ("FLOAT", {"default": 0.0}),
"safety": ("FLOAT", {"min": 0, "default": 0.9}),
"factor_min": ("FLOAT", {"min": 0, "default": 0.2}),
"factor_max": ("FLOAT", {"min": 0, "default": 10}),
"max_steps": ("INT", {"min": 1, "default": 2**31 - 1}),
"min_sigma": ("FLOAT", {"min": 0, "default": 1e-5}),
}
}
RETURN_TYPES = ("SAMPLER",)
FUNCTION = "get_sampler"
CATEGORY = "sampling/custom_sampling/samplers"
def get_sampler(self, **kwargs):
return (comfy.samplers.KSAMPLER(RungeKuttaSamplerImpl(**kwargs)),)
NODE_CLASS_MAPPINGS = {
"RungeKuttaSampler": RungeKuttaSampler,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"RungeKuttaSampler": "Runge-Kutta Sampler",
}
@@ -0,0 +1,277 @@
from typing import Any, Callable, Dict, Optional, Tuple
import torch
import torch.nn as nn
from torchode import status_codes
from torchode.problems import InitialValueProblem
from torchode.single_step_methods import StepResult
from torchode.step_size_controllers import PIDState, rms_norm
from torchode.terms import ODETerm
from torchode.typing import *
class PIDController(nn.Module):
"""A PID step size controller.
The formula for the dt scaling factor with PID control is taken from [1], Equation
(34).
References
----------
[1] Söderlind, G. (2003). Digital Filters in Adaptive Time-Stepping. ACM
Transactions on Mathematical Software, 29, 1–26.
"""
def __init__(
self,
atol: float,
rtol: float,
pcoeff: float,
icoeff: float,
dcoeff: float,
*,
term: Optional[ODETerm] = None,
norm: Callable[[DataTensor], NormTensor] = rms_norm,
force_monotonic_solve: Optional[bool] = True,
dt_min: Optional[float] = None,
dt_max: Optional[float] = None,
safety: float = 0.9,
factor_min: float = 0.2,
factor_max: float = 10.0,
):
super().__init__()
self.register_buffer("atol", torch.tensor(atol))
self.register_buffer("rtol", torch.tensor(rtol))
self.term = term
self.norm = norm
self.force_monotonic_solve = force_monotonic_solve
self.dt_min = dt_min
self.dt_max = dt_max
self.pcoeff = pcoeff
self.icoeff = icoeff
self.dcoeff = dcoeff
self.safety = safety
self.factor_min = factor_min
self.factor_max = factor_max
def dt_factor(self, state: PIDState, error_ratio: NormTensor):
"""Compute the growth factor of the timestep."""
# This is an instantiation of Equation (34) in the Söderlind paper where we have
# factored out the safety coefficient. I have not found a reference for dividing
# the PID coefficients by the order of the solver but DifferentialEquations.jl
# and diffrax both do it, so we do it too. Note that our error ratio is the
# reciprocal of Söderlind's error ratio (except for the safety factor).
# Therefore, the factor exponents have the opposite sign from the paper.
#
# Interesting thing from the introduction of that paper is that you work with p
# if you want per-step-error-control and p+1 if you want
# per-unit-step-error-control where p is the convergence order of the stepping
# method.
order = state.method_order
k_I, k_P, k_D = self.icoeff / order, self.pcoeff / order, self.dcoeff / order
factor1 = error_ratio ** (-(k_I + k_P + k_D))
factor2 = state.prev_error_ratio ** (k_P + 2 * k_D)
factor3 = state.prev_prev_error_ratio**-k_D
factor = self.safety * factor1 * factor2 * factor3
return torch.clamp(factor, min=self.factor_min, max=self.factor_max)
def initial_state(
self,
method_order: int,
problem: InitialValueProblem,
dt_min: Optional[TimeTensor],
dt_max: Optional[TimeTensor],
) -> PIDState:
return PIDState.default(
method_order=method_order,
batch_size=problem.batch_size,
dtype=problem.data_dtype,
device=problem.device,
dt_min=dt_min,
dt_max=dt_max,
)
@torch.jit.export
def merge_states(self, running: AcceptTensor, current: PIDState, previous: PIDState) -> PIDState:
return current.update_error_ratios(
torch.where(running, current.prev_error_ratio, previous.prev_error_ratio),
torch.where(running, current.prev_prev_error_ratio, previous.prev_prev_error_ratio),
)
def update_state(
self,
state: PIDState,
y0: DataTensor,
dt: TimeTensor,
error_ratio: Optional[NormTensor],
accept: Optional[AcceptTensor],
) -> PIDState:
if error_ratio is None:
return state.update_error_ratios(
prev_error_ratio=y0.new_ones(dt.shape),
prev_prev_error_ratio=state.prev_error_ratio,
)
else:
assert accept is not None
return state.update_error_ratios(
prev_error_ratio=torch.where(accept, error_ratio, state.prev_error_ratio),
prev_prev_error_ratio=torch.where(accept, state.prev_error_ratio, state.prev_prev_error_ratio),
)
################################################################################
# The following methods should be on AdaptiveStepSizeController if TorchScript #
# supports inheritance at some point #
################################################################################
@torch.jit.export
def init(
self,
term: Optional[ODETerm],
problem: InitialValueProblem,
method_order: int,
dt0: Optional[TimeTensor],
*,
stats: Dict[str, Any],
args: Any,
) -> Tuple[TimeTensor, PIDState, Optional[DataTensor]]:
if dt0 is None:
dt_max = (problem.t_end - problem.t_start).abs()
dt0, f0 = self._select_initial_step(
term,
problem.t_start,
problem.y0,
problem.time_direction,
dt_max,
method_order,
stats,
args,
)
else:
f0 = None
dt_min = self.dt_min
if dt_min is not None:
dt_min = torch.tensor(dt_min, dtype=problem.time_dtype, device=problem.device)
dt_max = self.dt_max
if dt_max is not None:
dt_max = torch.tensor(dt_max, dtype=problem.time_dtype, device=problem.device)
return dt0, self.initial_state(method_order, problem, dt_min, dt_max), f0
@torch.jit.export
def adapt_step_size(
self,
t0: TimeTensor,
dt: TimeTensor,
y0: DataTensor,
step_result: StepResult,
state: PIDState,
stats: Dict[str, Any],
) -> Tuple[AcceptTensor, TimeTensor, PIDState, Optional[StatusTensor]]:
y1, error_estimate = step_result.y, step_result.error_estimate
if error_estimate is None:
# If the stepping method could not provide an error estimate, we interpret
# this as an error estimate that gets the step accepted without changing the
# step size, i.e. as an error ratio of 1 (disregarding the safety factor).
return (
torch.ones_like(dt, dtype=torch.bool),
dt,
self.update_state(state, y0, dt, None, None),
None,
)
# Compute error ratio and decide on step acceptance
error_bounds = torch.add(self.atol, torch.maximum(y0.abs(), y1.abs()), alpha=self.rtol)
error = error_estimate.abs()
# We lower-bound the error ratio by some small number to avoid division by 0 in
# `dt_factor`.
error_ratio = torch.maximum(self.norm(error / error_bounds), state.almost_zero)
accept = error_ratio < 1.0
# Adapt the step size
dt_next = dt * self.dt_factor(state, error_ratio).to(dtype=dt.dtype)
# Check for infinities and NaN
status = torch.where(
torch.isfinite(error_ratio),
status_codes.SUCCESS,
status_codes.INFINITE_NORM,
)
# Enforce the minimum and maximum step size
dt_min = state.dt_min
dt_max = state.dt_max
if dt_min is not None or dt_max is not None:
dt_next = torch.clamp(dt_next, dt_min, dt_max)
return (
accept,
dt_next,
self.update_state(state, y0, dt, error_ratio, accept),
status,
)
def _select_initial_step(
self,
term: Optional[ODETerm],
t0: TimeTensor,
y0: DataTensor,
direction: torch.Tensor,
dt_max: TimeTensor,
convergence_order: int,
stats: Dict[str, Any],
args: Any,
) -> Tuple[TimeTensor, DataTensor]:
"""Empirically select a good initial step.
This is an adaptation of the algorithm described in [1]_. We changed it in such a
way that the tolerances apply to the norms instead of the components of `y`.
References
----------
.. [1] E. Hairer, S. P. Norsett G. Wanner, "Solving Ordinary Differential Equations
I: Nonstiff Problems", Sec. II.4, 2nd edition.
"""
if torch.jit.is_scripting() or term is None:
assert term is None, "The integration term is fixed for JIT compilation"
term = self.term
assert term is not None
norm = self.norm
f0 = term.vf(t0, y0, stats, args)
error_bounds = torch.add(self.atol, torch.abs(y0), alpha=self.rtol)
inv_scale = torch.reciprocal(error_bounds)
d0 = norm(y0 * inv_scale)
d1 = norm(f0 * inv_scale)
small_number = torch.tensor(1e-6, dtype=d0.dtype, device=d0.device)
dt0 = torch.where((d0 < 1e-5) | (d1 < 1e-5), small_number, 0.01 * d0 / d1)
# Ensure that we don't step out of the integration bounds
dt0 = torch.minimum(dt0, dt_max.to(dtype=y0.dtype))
y1 = torch.addcmul(y0, (direction * dt0)[:, None], f0)
f1 = term.vf(
torch.addcmul(t0, direction.to(dtype=t0.dtype), dt0.to(dtype=t0.dtype)),
y1,
stats,
args,
)
d2 = norm((f1 - f0) * inv_scale) / dt0
maxd1d2 = torch.maximum(d1, d2)
dt1 = torch.where(
maxd1d2 <= 1e-15,
torch.maximum(small_number, dt0 * 1e-3),
(0.01 / maxd1d2) ** (1.0 / convergence_order),
)
return (direction * torch.minimum(100 * dt0, dt1)).to(dtype=t0.dtype), f0
@@ -0,0 +1,61 @@
from typing import Any, Dict, Optional, Tuple
import torch
from torchode.problems import InitialValueProblem
from torchode.single_step_methods import StepResult
from torchode.step_size_controllers import StepSizeController
from torchode.terms import ODETerm
from torchode.typing import *
class ScheduledState:
def __init__(self, accept_all: AcceptTensor, step: int):
self.accept_all = accept_all
self.step = step
class ScheduledController(StepSizeController[ScheduledState]):
def __init__(self, sigmas: torch.Tensor):
super().__init__()
self.dt = [sigmas[i + 1] - sigmas[i] for i in range(len(sigmas) - 1)]
@torch.jit.export
def init(
self,
term: Optional[ODETerm],
problem: InitialValueProblem,
method_order: int,
dt0: Optional[TimeTensor],
*,
stats: Dict[str, Any],
args: Any,
):
assert dt0 is None
return (
self.dt[0],
ScheduledState(accept_all=torch.ones(problem.batch_size, device=problem.device, dtype=torch.bool), step=0),
None,
)
@torch.jit.export
def adapt_step_size(
self,
t0: TimeTensor,
dt: TimeTensor,
y0: DataTensor,
step_result: StepResult,
state: ScheduledState,
stats: Dict[str, Any],
) -> Tuple[AcceptTensor, TimeTensor, ScheduledState, Optional[StatusTensor]]:
state.step += 1
if state.step >= len(self.dt) - 1:
dt_next = -t0
else:
dt_next = self.dt[state.step]
return state.accept_all, dt_next, state, None
@torch.jit.export
def merge_states(self, running: AcceptTensor, current: ScheduledState, previous: ScheduledState) -> ScheduledState:
return current
+16
View File
@@ -0,0 +1,16 @@
[project]
name = "ComfyUI-RK-Sampler"
description = "Batched Runge-Kutta Samplers for ComfyUI"
version = "0.0.1"
license = "LICENSE"
dependencies = [
"torchode"
]
[project.urls]
Repository = "https://github.com/wootwootwootwoot/ComfyUI-RK-Sampler"
[tool.comfy]
PublisherId = "wootwootwootwoot"
DisplayName = "ComfyUI-RK-Sampler"
Icon = ""
+1
View File
@@ -0,0 +1 @@
torchode