358 lines
14 KiB
Python
358 lines
14 KiB
Python
import matplotlib.pyplot as plt
|
|
import numpy as np
|
|
import os
|
|
from ovf.core.basis.MultiPortOrthonormalBasis import MultiPortOrthonormalBasis
|
|
from ovf.core.utils import generate_starting_poles
|
|
import json
|
|
import pickle
|
|
|
|
class VFManager():
|
|
def __init__(
|
|
self,
|
|
npoles_cplx,
|
|
freqs,
|
|
H,
|
|
model=MultiPortOrthonormalBasis,
|
|
iterations:int=5,
|
|
fit_constant:bool=True,
|
|
fit_proportional:bool=False,
|
|
dc_enforce:bool=False,
|
|
passivity_enforce:bool=True,
|
|
verbose:bool=True
|
|
):
|
|
|
|
|
|
self.freqs=freqs
|
|
self.H=H
|
|
self.iterations=iterations
|
|
self.fit_constant=fit_constant
|
|
self.fit_proportional=fit_proportional
|
|
self.dc_enforce=dc_enforce
|
|
self.passivity_enforce=passivity_enforce
|
|
self.verbose=verbose
|
|
|
|
self.nports = H.shape[1]
|
|
self.npoles_cplx = npoles_cplx
|
|
|
|
self.least_squares_condition = []
|
|
self.least_squares_row_condition = []
|
|
self.least_squares_col_condition = []
|
|
self.least_squares_rms_error = []
|
|
self.eigenval_condition = []
|
|
self.eigenval_row_condition = []
|
|
self.eigenval_col_condition = []
|
|
self.eigenval_rms_error = []
|
|
|
|
self.model_instance:MultiPortOrthonormalBasis|None = None
|
|
self.model_responses_freqs = None
|
|
self.model_responses_H = None
|
|
|
|
self.model=model
|
|
|
|
@property
|
|
def residuals(self):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
return self.model_instance.residuals
|
|
|
|
@property
|
|
def constant(self):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
return self.model_instance.D
|
|
|
|
@property
|
|
def proportional(self):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
return self.model_instance.proportional
|
|
|
|
@property
|
|
def poles(self):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
return list(self.model_instance.poles)
|
|
|
|
@property
|
|
def rms_error(self):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
return self.eigenval_rms_error[-1]
|
|
|
|
@property
|
|
def condition(self):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
return self.eigenval_condition[-1]
|
|
|
|
@property
|
|
def row_condition(self):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
return self.eigenval_row_condition[-1]
|
|
|
|
@property
|
|
def col_condition(self):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
return self.eigenval_col_condition[-1]
|
|
|
|
|
|
@classmethod
|
|
def load(cls,dirname):
|
|
instance = cls._load_npz(f"{dirname}/model.npz")
|
|
instance._load_model_instance(f"{dirname}/model_instance.pkl")
|
|
return instance
|
|
|
|
def write(self,dirname):
|
|
os.makedirs(dirname, exist_ok=True)
|
|
self._write_npz(f"{dirname}/model.npz")
|
|
self._write_model_instance(f"{dirname}/model_instance.pkl")
|
|
|
|
def _write_npz(self,filename):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
np.savez(filename,
|
|
model_name=self.model.__name__,
|
|
poles=self.model_instance.poles,
|
|
freqs=self.freqs,
|
|
H=self.H,
|
|
iterations=self.iterations,
|
|
verbose=self.verbose,
|
|
nports=self.nports,
|
|
npoles_cplx=self.npoles_cplx,
|
|
least_squares_condition=self.least_squares_condition,
|
|
least_squares_row_condition=self.least_squares_row_condition,
|
|
least_squares_col_condition=self.least_squares_col_condition,
|
|
least_squares_rms_error=self.least_squares_rms_error,
|
|
eigenval_condition=self.eigenval_condition,
|
|
eigenval_row_condition=self.eigenval_row_condition,
|
|
eigenval_col_condition=self.eigenval_col_condition,
|
|
eigenval_rms_error=self.eigenval_rms_error,
|
|
fit_constant=self.fit_constant,
|
|
fit_proportional=self.fit_proportional,
|
|
dc_enforce=self.dc_enforce,
|
|
passivity_enforce=self.passivity_enforce,
|
|
denominator=self.model_instance.denominator,
|
|
allow_pickle=True
|
|
)
|
|
if self.verbose:
|
|
print(f"Model parameters saved to {filename}")
|
|
|
|
def _write_model_instance(self,filename):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
with open(filename,"wb") as f:
|
|
pickle.dump(self.model_instance,f)
|
|
if self.verbose:
|
|
print(f"Model instance saved to {filename}")
|
|
|
|
def _load_model_instance(self,filename):
|
|
with open(filename,"rb") as f:
|
|
self.model_instance = pickle.load(f)
|
|
if self.verbose:
|
|
print(f"Model instance loaded from {filename}")
|
|
return self.model_instance
|
|
|
|
@classmethod
|
|
def _load_npz(cls,filename):
|
|
instance = cls(
|
|
npoles_cplx=1, # 临时值,稍后会被覆盖
|
|
freqs=np.array([]), # 临时值,稍后会被覆盖
|
|
H=np.array([[]]), # 临时值,稍后会被覆盖
|
|
model=MultiPortOrthonormalBasis, # 临时值,稍后会被覆盖
|
|
iterations=1, # 临时值,稍后会被覆盖
|
|
verbose=False # 临时值,稍后会被覆盖
|
|
)
|
|
data = np.load(filename,allow_pickle=True)
|
|
instance.freqs = data["freqs"]
|
|
instance.H = data["H"]
|
|
instance.iterations = int(data["iterations"])
|
|
instance.verbose = bool(data["verbose"])
|
|
instance.nports = int(data["nports"])
|
|
instance.npoles_cplx = int(data["npoles_cplx"])
|
|
instance.least_squares_condition = data["least_squares_condition"].tolist()
|
|
instance.least_squares_row_condition = data["least_squares_row_condition"].tolist()
|
|
instance.least_squares_col_condition = data["least_squares_col_condition"].tolist()
|
|
instance.least_squares_rms_error = data["least_squares_rms_error"].tolist()
|
|
instance.eigenval_condition = data["eigenval_condition"].tolist()
|
|
instance.eigenval_row_condition = data["eigenval_row_condition"].tolist()
|
|
instance.eigenval_col_condition = data["eigenval_col_condition"].tolist()
|
|
instance.eigenval_rms_error = data["eigenval_rms_error"].tolist()
|
|
instance.fit_constant = bool(data["fit_constant"])
|
|
instance.fit_proportional = bool(data["fit_proportional"])
|
|
instance.dc_enforce = bool(data["dc_enforce"])
|
|
instance.passivity_enforce = bool(data["passivity_enforce"])
|
|
poles = data["poles"]
|
|
denominator = data["denominator"]
|
|
instance.model = globals()[data["model_name"].item()]
|
|
if instance.verbose:
|
|
print(f"Model parameters loaded from {filename}")
|
|
return instance
|
|
|
|
|
|
def fit(self):
|
|
self.levi()
|
|
self.model_instance = self.sk_iteration()
|
|
return self.model
|
|
|
|
def levi(self):
|
|
self._poles = generate_starting_poles(self.npoles_cplx,beta_min=max(self.freqs[0],1e4),beta_max=self.freqs[-1]*1.1,verbose=self.verbose)
|
|
self.model_instance=self.model(
|
|
H=self.H,
|
|
freqs=self.freqs,
|
|
pre_poles=self._poles,
|
|
fit_constant=self.fit_constant,
|
|
fit_proportional=self.fit_proportional,
|
|
dc_enforce=self.dc_enforce,
|
|
passivity_enforce=self.passivity_enforce
|
|
)
|
|
return self.model_instance
|
|
|
|
def sk_iteration(self):
|
|
for i in range(self.iterations):
|
|
assert self.model_instance is not None ,"Please run levi() first."
|
|
self._poles = self.model_instance.poles
|
|
self.model_instance = self.model(
|
|
H=self.H,
|
|
freqs=self.freqs,
|
|
pre_poles=self._poles,
|
|
fit_constant=self.fit_constant,
|
|
fit_proportional=self.fit_proportional,
|
|
dc_enforce=self.dc_enforce,
|
|
passivity_enforce=self.passivity_enforce
|
|
)
|
|
if self.verbose:
|
|
print(f"Iteration {i+1}/{self.iterations}")
|
|
print("A:",self.model_instance.A)
|
|
print("B:",self.model_instance.B)
|
|
print("C:",self.model_instance.C)
|
|
print("D:",self.model_instance.D)
|
|
print("poles:",self.model_instance.poles)
|
|
print("denominator:",self.model_instance.denominator)
|
|
|
|
least_squares_condition,\
|
|
least_squares_row_condition,\
|
|
least_squares_col_condition,\
|
|
least_squares_rms_error = self.model_instance.least_squares_metric
|
|
eigenval_condition,\
|
|
eigenval_row_condition,\
|
|
eigenval_col_condition,\
|
|
eigenval_rms_error = self.model_instance.eigen_metric
|
|
self.least_squares_condition.append(least_squares_condition)
|
|
self.least_squares_row_condition.append(least_squares_row_condition)
|
|
self.least_squares_col_condition.append(least_squares_col_condition)
|
|
self.least_squares_rms_error.append(least_squares_rms_error)
|
|
self.eigenval_condition.append(eigenval_condition)
|
|
self.eigenval_row_condition.append(eigenval_row_condition)
|
|
self.eigenval_col_condition.append(eigenval_col_condition)
|
|
self.eigenval_rms_error.append(eigenval_rms_error)
|
|
return self.model_instance
|
|
|
|
def plot_metrics(self,show:bool=True,save_path=None):
|
|
plt.figure(figsize=(16, 12))
|
|
plt.subplot(4, 2, 1)
|
|
plt.plot(
|
|
range(1, len(self.least_squares_condition) + 1),
|
|
self.least_squares_condition,
|
|
label='Least Squares Condition'
|
|
)
|
|
plt.legend()
|
|
plt.subplot(4, 2, 2)
|
|
plt.plot(
|
|
range(1, len(self.least_squares_row_condition) + 1),
|
|
self.least_squares_row_condition,
|
|
label='Least Squares Row Condition'
|
|
)
|
|
plt.legend()
|
|
|
|
plt.subplot(4, 2, 3)
|
|
plt.plot(
|
|
range(1, len(self.least_squares_col_condition) + 1),
|
|
self.least_squares_col_condition,
|
|
label='Least Squares Col Condition'
|
|
)
|
|
plt.legend()
|
|
|
|
plt.subplot(4, 2, 4)
|
|
plt.plot(
|
|
range(1, len(self.least_squares_rms_error) + 1),
|
|
self.least_squares_rms_error,
|
|
label='Least Squares RMS Error'
|
|
)
|
|
plt.legend()
|
|
|
|
plt.subplot(4, 2, 5)
|
|
plt.plot(
|
|
range(1, len(self.eigenval_condition) + 1),
|
|
self.eigenval_condition,
|
|
label='Eigenvalue Condition'
|
|
)
|
|
plt.legend()
|
|
|
|
plt.subplot(4, 2, 6)
|
|
plt.plot(
|
|
range(1, len(self.eigenval_row_condition) + 1),
|
|
self.eigenval_row_condition,
|
|
label='Eigenvalue Row Condition'
|
|
)
|
|
plt.legend()
|
|
|
|
plt.subplot(4, 2, 7)
|
|
plt.plot(
|
|
range(1, len(self.eigenval_col_condition) + 1),
|
|
self.eigenval_col_condition,
|
|
label='Eigenvalue Col Condition'
|
|
)
|
|
plt.legend()
|
|
|
|
plt.subplot(4, 2, 8)
|
|
plt.plot(
|
|
range(1, len(self.eigenval_rms_error) + 1),
|
|
self.eigenval_rms_error,
|
|
label='Eigenvalue RMS Error'
|
|
)
|
|
plt.legend()
|
|
|
|
if show:
|
|
plt.show()
|
|
if save_path is not None:
|
|
if self.verbose:
|
|
print(f"Saving metrics plot to {save_path}/fitting_metrics.png")
|
|
os.makedirs(save_path, exist_ok=True)
|
|
plt.savefig(f"{save_path}/fitting_metrics.png")
|
|
|
|
def plot_model_responses(self,show:bool=True,save_path=None):
|
|
assert self.model_responses_freqs is not None and self.model_responses_H is not None, "Please run get_model_responses() first."
|
|
for i in range(self.nports):
|
|
for j in range(self.nports):
|
|
plt.figure(figsize=(12, 6))
|
|
plt.subplot(2, 2, 1)
|
|
plt.plot(self.freqs, np.abs(self.H[:,i,j]), 'o', ms=4, color='red', label='Input Samples')
|
|
plt.plot(self.model_responses_freqs, np.abs(self.model_responses_H[:,i,j]), '-', lw=2, color='k', label='Fit')
|
|
plt.title(f"Response i={i+1}, j={j+1}")
|
|
plt.ylabel("Magnitude")
|
|
plt.legend(loc="best")
|
|
plt.subplot(2, 2, 2)
|
|
plt.plot(self.freqs, np.angle(self.H[:,i,j],deg=True), 'o', ms=4, color='red', label='Input Samples')
|
|
plt.plot(self.model_responses_freqs, np.angle(self.model_responses_H[:,i,j],deg=True), '-', lw=2, color='k', label='Fit')
|
|
plt.title(f"Response i={i+1}, j={j+1}")
|
|
plt.ylabel("Phase (deg)")
|
|
plt.legend(loc="best")
|
|
plt.tight_layout()
|
|
plt.subplot(2, 2, 3)
|
|
plt.plot(self.freqs, np.real(self.H[:,i,j]), 'o', ms=4, color='red', label='Input Samples')
|
|
plt.plot(self.model_responses_freqs, np.real(self.model_responses_H[:,i,j]), '-', lw=2, color='k', label='Fit')
|
|
plt.title(f"Response i={i+1}, j={j+1}")
|
|
plt.ylabel("Real Part")
|
|
plt.legend(loc="best")
|
|
plt.subplot(2, 2, 4)
|
|
plt.plot(self.freqs, np.imag(self.H[:,i,j]), 'o', ms=4, color='red', label='Input Samples')
|
|
plt.plot(self.model_responses_freqs, np.imag(self.model_responses_H[:,i,j]), '-', lw=2, color='k', label='Fit')
|
|
plt.title(f"Response i={i+1}, j={j+1}")
|
|
plt.ylabel("Imag Part")
|
|
plt.legend(loc="best")
|
|
plt.tight_layout()
|
|
if show:
|
|
plt.show()
|
|
if save_path is not None:
|
|
if self.verbose:
|
|
print(f"Saving response plot for port {i+1},{j+1} to {save_path}/response_{i+1}_{j+1}.png")
|
|
os.makedirs(save_path, exist_ok=True)
|
|
plt.savefig(f"{save_path}/response_{i+1}_{j+1}.png")
|
|
|
|
def get_model_responses(self,freqs):
|
|
assert self.model_instance is not None ,"Please run levi() and sk_iteration() first."
|
|
self.model_responses_freqs = freqs
|
|
self.model_responses_H = self.model_instance.get_model_responses(freqs)
|
|
return self.model_responses_H |