3131from .. import Unit
3232from .design_tools import MESH
3333from numba import float64 , int8 , types , njit
34+ from typing import NamedTuple
3435from thermosteam import (
3536 equilibrium , VariableNode ,
3637)
4344 'PhasePartition' ,
4445)
4546
47+ class IterationResult (NamedTuple ):
48+ t : float #: Step size
49+ x : np .ndarray #: Point
50+ r : float #: Residual
51+
52+
4653# %% Equation-oriented tools
4754
4855@jitdata
@@ -1855,8 +1862,10 @@ class MultiStageEquilibrium(Unit):
18551862 """
18561863 _N_ins = 2
18571864 _N_outs = 2
1865+ _line_search = False # Experimental feature
1866+ dynamic_learning_rate = 0.1
18581867 inside_maxiter = 100
1859- default_max_attempts = 10
1868+ default_max_attempts = 5
18601869 default_maxiter = 100
18611870 default_optimize_result = 'trust-region' , 50
18621871 default_tolerance = 1e-5
@@ -2526,23 +2535,20 @@ def _run(self):
25262535 self .attempt = 0
25272536 self .bottom_flows = None
25282537 self ._mean_residual = np .inf
2529- empty = flx . LineSearchResult (None , None , np .inf )
2530- residual_record = 5 * [empty ]
2531- residual_record [0 ] = flx . LineSearchResult (1 , x , self ._objective (x ))
2532- self ._residual_record = deque (residual_record )
2538+ self . _best_result = empty = IterationResult (None , None , np .inf )
2539+ record = 5 * [empty ]
2540+ record [0 ] = IterationResult (1 , x , self ._objective (x ))
2541+ self ._residual_record = record = deque (record )
25332542 if self .vle_decomposition is None :
25342543 self .default_vle_decomposition ()
25352544 self .iter = 0
25362545 for n in range (self .max_attempts ):
25372546 self .attempt = n
25382547 try : x = solver (f , self ._get_point (), ** options )
2539- except :
2540- residual = self ._mean_residual
2548+ except :
25412549 self ._mean_residual = np .inf
2542- self ._residual_record [1 ] = empty
25432550 try : x = solver (self ._sequential_iter , self ._get_point (), ** options )
2544- except :
2545- if residual >= self ._mean_residual : break
2551+ except : break
25462552 else :
25472553 self ._set_point (x )
25482554 break
@@ -2553,7 +2559,7 @@ def _run(self):
25532559 original = self .method , self .maxiter
25542560 self .method , self .maxiter = optimize_result
25552561 try :
2556- self ._simultaneous_correction (x )
2562+ self ._simultaneous_correction (self . _get_point () )
25572563 finally :
25582564 self .method , self .maxiter = original
25592565 elif algorithm == 'simultaneous correction' :
@@ -2565,30 +2571,31 @@ def _run(self):
25652571
25662572 def _simultaneous_correction (self , x ):
25672573 shape = x .shape
2568- x = x .flatten ()
25692574 f = lambda x : self ._residuals (x .reshape (shape )).flatten ()
25702575 jac = lambda x : MESH .create_block_tridiagonal_matrix (* self ._jacobian (x .reshape (shape )))
25712576 # x, *self._simultaneous_correction_info = leastsq(f, x, Dfun=jac, full_output=True, maxfev=self.maxiter, xtol=self.tolerance)
25722577 # f = lambda x: self._objective(x.reshape(shape))
25732578 # self._simultaneous_correction_info = res = minimize_ipopt(f, x, jac=jac)
25742579 # x = res.x
2575- try : x , * self ._simultaneous_correction_info = fsolve (f , x , fprime = jac , full_output = True , maxfev = self .maxiter , xtol = self .tolerance )
2580+ try :
2581+ x , * self ._simultaneous_correction_info = fsolve (
2582+ f , x .flatten (), fprime = jac , full_output = True ,
2583+ maxfev = self .maxiter , xtol = self .tolerance
2584+ )
25762585 except : pass
25772586 else :
25782587 x [x < 0 ] = 0
25792588 x = x .reshape (shape )
25802589 r = self ._objective (x )
25812590 try :
2582- record = self ._residual_record
2591+ result = self ._best_result
25832592 except :
25842593 self ._set_point (x )
2585- self .update_mass_balance ()
25862594 else :
2587- index = np .argmin ([i .f for i in record ])
2588- result = record [index ]
2589- if r < result .f :
2590- self ._set_point (x )
2591- self .update_mass_balance ()
2595+ self ._set_point (x )
2596+ self .update_mass_balance ()
2597+ r = self ._objective (self ._get_point ())
2598+ if result .r < r : self ._set_point (result .x )
25922599
25932600 def _phenomena_iter (self , x0 ):
25942601 self .iter += 1
@@ -2601,7 +2608,6 @@ def _phenomena_iter(self, x0):
26012608 for i in self .stages : i ._update_separation_factors ()
26022609 separation_factors = np .array ([i .S for i in self ._S_stages ])
26032610 self .update_flow_rates (separation_factors , update_B = True )
2604- return self ._line_search (x0 , self ._get_point ())
26052611 elif decomp == 'sum rates' :
26062612 self .update_pseudo_vle ()
26072613 separation_factors = np .array ([i .S for i in self ._S_stages ])
@@ -2614,28 +2620,33 @@ def _phenomena_iter(self, x0):
26142620 self .update_energy_balance_temperatures ()
26152621 else :
26162622 raise RuntimeError ('unknown equilibrium phenomena' )
2617- return self ._get_point ()
2618-
2619- def _line_search (self , x0 , x1 ):
2620- correction = x1 - x0
2623+ return self ._new_point ()
2624+
2625+ def _new_point (self ):
26212626 record = self ._residual_record
2622- t0 , _ , r0 = record [0 ]
2623- result = flx .inexact_line_search (
2624- self ._objective , x0 , correction ,
2625- fx = r0 , t0 = 0.05 , t1 = 1.2 , tguess = t0
2626- )
2627+ t0 , x0 , r0 = record [0 ]
2628+ x1 = self ._get_point ()
2629+ correction = x1 - x0
2630+ if self ._line_search :
2631+ result = flx .inexact_line_search (
2632+ self ._objective , x0 , correction ,
2633+ fx = r0 , t0 = 0.8 , t1 = 1.2 , tguess = t0
2634+ )
2635+ else :
2636+ result = IterationResult (1 , x1 , self ._objective (x1 ))
26272637 x1 = result .x
26282638 x1 [x1 < 0 ] = 0
26292639 record .rotate ()
26302640 record [0 ] = result
2631- residuals = np .array ([i .f for i in record ])
2641+ residuals = np .array ([i .r for i in record ])
26322642 mean = np .mean (residuals )
26332643 if mean > self ._mean_residual :
2634- index = np .argmin (residuals )
2635- record [0 ] = result = record [index ]
2644+ record [0 ] = result = self ._best_result
26362645 self ._set_point (result .x )
26372646 raise RuntimeError ('residual error is oscillating' )
26382647 else :
2648+ if self ._best_result .r > result .r :
2649+ self ._best_result = result
26392650 self ._mean_residual = mean
26402651 return x1
26412652
@@ -2644,7 +2655,7 @@ def _sequential_iter(self, x0):
26442655 self ._set_point (x0 )
26452656 for i in self .stages : i ._run ()
26462657 for i in reversed (self .stages ): i ._run ()
2647- return self ._line_search ( x0 , self . _get_point () )
2658+ return self ._new_point ( )
26482659
26492660 def _iter (self , x0 ):
26502661 algorithm = self .algorithm
0 commit comments