Initial commit
This commit is contained in:
@@ -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
@@ -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,
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
{
|
||||
"RungeKuttaSampler": ""
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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]
|
||||
)
|
||||
@@ -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]
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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 = ""
|
||||
@@ -0,0 +1 @@
|
||||
torchode
|
||||
Reference in New Issue
Block a user