From be0e6fd0bbe36ed1e369977bbe5a598326021eff Mon Sep 17 00:00:00 2001 From: rossiyareich Date: Sun, 21 Jul 2024 23:00:34 +0700 Subject: [PATCH] Initial commit --- .gitignore | 2 + __init__.py | 16 + node_list.json | 3 + nodes/__init__.py | 0 nodes/methods/__init__.py | 0 nodes/methods/ae_bosh3.py | 29 ++ nodes/methods/ae_cash_karp5.py | 43 +++ nodes/methods/ae_dopri5.py | 58 ++++ nodes/methods/ae_dopri8.py | 264 ++++++++++++++ nodes/methods/ae_fehlberg2.py | 29 ++ nodes/methods/ae_fehlberg5.py | 43 +++ nodes/methods/ae_heun_euler2.py | 29 ++ nodes/methods/ae_midpoint2.py | 29 ++ nodes/methods/ae_ralston2.py | 29 ++ nodes/methods/ae_tsit5.py | 134 ++++++++ nodes/methods/fe_euler1.py | 28 ++ nodes/methods/fe_heun3.py | 28 ++ nodes/methods/fe_kutta3.py | 28 ++ nodes/methods/fe_kutta4.py | 28 ++ nodes/methods/fe_kutta_38th4.py | 28 ++ nodes/methods/fe_ralston3.py | 28 ++ nodes/methods/fe_ralston4.py | 41 +++ nodes/methods/fe_ssprk3.py | 28 ++ nodes/methods/fe_wray3.py | 28 ++ nodes/methods/runge_kutta.py | 249 ++++++++++++++ nodes/nodes_rk_sampler.py | 321 ++++++++++++++++++ nodes/step_size_controllers/__init__.py | 0 nodes/step_size_controllers/pid_controller.py | 277 +++++++++++++++ .../scheduled_controller.py | 61 ++++ pyproject.toml | 16 + requirements.txt | 1 + 31 files changed, 1898 insertions(+) create mode 100644 __init__.py create mode 100644 node_list.json create mode 100644 nodes/__init__.py create mode 100644 nodes/methods/__init__.py create mode 100644 nodes/methods/ae_bosh3.py create mode 100644 nodes/methods/ae_cash_karp5.py create mode 100644 nodes/methods/ae_dopri5.py create mode 100644 nodes/methods/ae_dopri8.py create mode 100644 nodes/methods/ae_fehlberg2.py create mode 100644 nodes/methods/ae_fehlberg5.py create mode 100644 nodes/methods/ae_heun_euler2.py create mode 100644 nodes/methods/ae_midpoint2.py create mode 100644 nodes/methods/ae_ralston2.py create mode 100644 nodes/methods/ae_tsit5.py create mode 100644 nodes/methods/fe_euler1.py create mode 100644 nodes/methods/fe_heun3.py create mode 100644 nodes/methods/fe_kutta3.py create mode 100644 nodes/methods/fe_kutta4.py create mode 100644 nodes/methods/fe_kutta_38th4.py create mode 100644 nodes/methods/fe_ralston3.py create mode 100644 nodes/methods/fe_ralston4.py create mode 100644 nodes/methods/fe_ssprk3.py create mode 100644 nodes/methods/fe_wray3.py create mode 100644 nodes/methods/runge_kutta.py create mode 100644 nodes/nodes_rk_sampler.py create mode 100644 nodes/step_size_controllers/__init__.py create mode 100644 nodes/step_size_controllers/pid_controller.py create mode 100644 nodes/step_size_controllers/scheduled_controller.py create mode 100644 pyproject.toml create mode 100644 requirements.txt diff --git a/.gitignore b/.gitignore index 82f9275..a2690e6 100644 --- a/.gitignore +++ b/.gitignore @@ -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/ \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..6391968 --- /dev/null +++ b/__init__.py @@ -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, +} diff --git a/node_list.json b/node_list.json new file mode 100644 index 0000000..823f15c --- /dev/null +++ b/node_list.json @@ -0,0 +1,3 @@ +{ + "RungeKuttaSampler": "" +} \ No newline at end of file diff --git a/nodes/__init__.py b/nodes/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/nodes/methods/__init__.py b/nodes/methods/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/nodes/methods/ae_bosh3.py b/nodes/methods/ae_bosh3.py new file mode 100644 index 0000000..0b79191 --- /dev/null +++ b/nodes/methods/ae_bosh3.py @@ -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) diff --git a/nodes/methods/ae_cash_karp5.py b/nodes/methods/ae_cash_karp5.py new file mode 100644 index 0000000..654eb21 --- /dev/null +++ b/nodes/methods/ae_cash_karp5.py @@ -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) diff --git a/nodes/methods/ae_dopri5.py b/nodes/methods/ae_dopri5.py new file mode 100644 index 0000000..5408e84 --- /dev/null +++ b/nodes/methods/ae_dopri5.py @@ -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] + ) diff --git a/nodes/methods/ae_dopri8.py b/nodes/methods/ae_dopri8.py new file mode 100644 index 0000000..27a3270 --- /dev/null +++ b/nodes/methods/ae_dopri8.py @@ -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] + ) diff --git a/nodes/methods/ae_fehlberg2.py b/nodes/methods/ae_fehlberg2.py new file mode 100644 index 0000000..070fb47 --- /dev/null +++ b/nodes/methods/ae_fehlberg2.py @@ -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) diff --git a/nodes/methods/ae_fehlberg5.py b/nodes/methods/ae_fehlberg5.py new file mode 100644 index 0000000..b2a7aed --- /dev/null +++ b/nodes/methods/ae_fehlberg5.py @@ -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) diff --git a/nodes/methods/ae_heun_euler2.py b/nodes/methods/ae_heun_euler2.py new file mode 100644 index 0000000..298a6f5 --- /dev/null +++ b/nodes/methods/ae_heun_euler2.py @@ -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) diff --git a/nodes/methods/ae_midpoint2.py b/nodes/methods/ae_midpoint2.py new file mode 100644 index 0000000..3c6144c --- /dev/null +++ b/nodes/methods/ae_midpoint2.py @@ -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) diff --git a/nodes/methods/ae_ralston2.py b/nodes/methods/ae_ralston2.py new file mode 100644 index 0000000..94d7b56 --- /dev/null +++ b/nodes/methods/ae_ralston2.py @@ -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) diff --git a/nodes/methods/ae_tsit5.py b/nodes/methods/ae_tsit5.py new file mode 100644 index 0000000..b80e2d1 --- /dev/null +++ b/nodes/methods/ae_tsit5.py @@ -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) diff --git a/nodes/methods/fe_euler1.py b/nodes/methods/fe_euler1.py new file mode 100644 index 0000000..78d9f53 --- /dev/null +++ b/nodes/methods/fe_euler1.py @@ -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) diff --git a/nodes/methods/fe_heun3.py b/nodes/methods/fe_heun3.py new file mode 100644 index 0000000..6bb20c4 --- /dev/null +++ b/nodes/methods/fe_heun3.py @@ -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) diff --git a/nodes/methods/fe_kutta3.py b/nodes/methods/fe_kutta3.py new file mode 100644 index 0000000..2eaceb3 --- /dev/null +++ b/nodes/methods/fe_kutta3.py @@ -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) diff --git a/nodes/methods/fe_kutta4.py b/nodes/methods/fe_kutta4.py new file mode 100644 index 0000000..a86e7ff --- /dev/null +++ b/nodes/methods/fe_kutta4.py @@ -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) diff --git a/nodes/methods/fe_kutta_38th4.py b/nodes/methods/fe_kutta_38th4.py new file mode 100644 index 0000000..8363fe2 --- /dev/null +++ b/nodes/methods/fe_kutta_38th4.py @@ -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) diff --git a/nodes/methods/fe_ralston3.py b/nodes/methods/fe_ralston3.py new file mode 100644 index 0000000..7670e04 --- /dev/null +++ b/nodes/methods/fe_ralston3.py @@ -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) diff --git a/nodes/methods/fe_ralston4.py b/nodes/methods/fe_ralston4.py new file mode 100644 index 0000000..ba9493c --- /dev/null +++ b/nodes/methods/fe_ralston4.py @@ -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) diff --git a/nodes/methods/fe_ssprk3.py b/nodes/methods/fe_ssprk3.py new file mode 100644 index 0000000..c1e151a --- /dev/null +++ b/nodes/methods/fe_ssprk3.py @@ -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) diff --git a/nodes/methods/fe_wray3.py b/nodes/methods/fe_wray3.py new file mode 100644 index 0000000..029da79 --- /dev/null +++ b/nodes/methods/fe_wray3.py @@ -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) diff --git a/nodes/methods/runge_kutta.py b/nodes/methods/runge_kutta.py new file mode 100644 index 0000000..d72060b --- /dev/null +++ b/nodes/methods/runge_kutta.py @@ -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() diff --git a/nodes/nodes_rk_sampler.py b/nodes/nodes_rk_sampler.py new file mode 100644 index 0000000..5d2f399 --- /dev/null +++ b/nodes/nodes_rk_sampler.py @@ -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", +} diff --git a/nodes/step_size_controllers/__init__.py b/nodes/step_size_controllers/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/nodes/step_size_controllers/pid_controller.py b/nodes/step_size_controllers/pid_controller.py new file mode 100644 index 0000000..fceb2db --- /dev/null +++ b/nodes/step_size_controllers/pid_controller.py @@ -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 diff --git a/nodes/step_size_controllers/scheduled_controller.py b/nodes/step_size_controllers/scheduled_controller.py new file mode 100644 index 0000000..3878198 --- /dev/null +++ b/nodes/step_size_controllers/scheduled_controller.py @@ -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 diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..7f7bcd5 --- /dev/null +++ b/pyproject.toml @@ -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 = "" diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..109fd0a --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +torchode