@@ -841,3 +841,44 @@ def sample_heunpp2(model, x, sigmas, extra_args=None, callback=None, disable=Non
841841 d_prime = w1 * d + w2 * d_2 + w3 * d_3
842842 x = x + d_prime * dt
843843 return x
844+
845+
846+ #From https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py
847+ #under Apache 2 license
848+ def sample_ipndm (model , x , sigmas , extra_args = None , callback = None , disable = None , max_order = 4 ):
849+ extra_args = {} if extra_args is None else extra_args
850+ s_in = x .new_ones ([x .shape [0 ]])
851+
852+ x_next = x
853+
854+ buffer_model = []
855+ for i in trange (len (sigmas ) - 1 , disable = disable ):
856+ t_cur = sigmas [i ]
857+ t_next = sigmas [i + 1 ]
858+
859+ x_cur = x_next
860+
861+ denoised = model (x_cur , t_cur * s_in , ** extra_args )
862+ if callback is not None :
863+ callback ({'x' : x , 'i' : i , 'sigma' : sigmas [i ], 'sigma_hat' : sigmas [i ], 'denoised' : denoised })
864+
865+ d_cur = (x_cur - denoised ) / t_cur
866+
867+ order = min (max_order , i + 1 )
868+ if order == 1 : # First Euler step.
869+ x_next = x_cur + (t_next - t_cur ) * d_cur
870+ elif order == 2 : # Use one history point.
871+ x_next = x_cur + (t_next - t_cur ) * (3 * d_cur - buffer_model [- 1 ]) / 2
872+ elif order == 3 : # Use two history points.
873+ x_next = x_cur + (t_next - t_cur ) * (23 * d_cur - 16 * buffer_model [- 1 ] + 5 * buffer_model [- 2 ]) / 12
874+ elif order == 4 : # Use three history points.
875+ x_next = x_cur + (t_next - t_cur ) * (55 * d_cur - 59 * buffer_model [- 1 ] + 37 * buffer_model [- 2 ] - 9 * buffer_model [- 3 ]) / 24
876+
877+ if len (buffer_model ) == max_order - 1 :
878+ for k in range (max_order - 2 ):
879+ buffer_model [k ] = buffer_model [k + 1 ]
880+ buffer_model [- 1 ] = d_cur
881+ else :
882+ buffer_model .append (d_cur )
883+
884+ return x_next
0 commit comments