mirror of
https://github.com/wassname/simpeg.git
synced 2026-08-19 12:40:13 +08:00
Jvec working for MT1D, Jtvec getting close
This commit is contained in:
+38
-39
@@ -2,7 +2,8 @@ from simpegEM.FDEM import BaseFDEMProblem
|
||||
from SurveyMT import SurveyMT
|
||||
from DataMT import DataMT
|
||||
from FieldsMT import FieldsMT
|
||||
from SimPEG import SolverLU as SimpegSolver
|
||||
from SimPEG import SolverLU as SimpegSolver, mkvc
|
||||
import numpy as np
|
||||
|
||||
class BaseMTProblem(BaseFDEMProblem):
|
||||
|
||||
@@ -24,18 +25,32 @@ class BaseMTProblem(BaseFDEMProblem):
|
||||
# Might need to add more stuff here.
|
||||
|
||||
def Jvec(self, m, v, f=None):
|
||||
"""
|
||||
Function to calculate the data sensitivities dD/dm times a vector.
|
||||
|
||||
:param numpy.ndarray (nC, 1) - conductive model
|
||||
:param numpy.ndarray (nC, 1) - random vector
|
||||
:param MTfields object (optional) - MT fields object, if not given it is calculated
|
||||
:rtype: MTdata object
|
||||
:return: Data sensitivities wrt m
|
||||
"""
|
||||
|
||||
# Calculate the fields
|
||||
if f is None:
|
||||
f = self.fields(m)
|
||||
|
||||
# Set current model
|
||||
self.curModel = m
|
||||
|
||||
# Initiate the Jv object
|
||||
Jv = self.dataPair(self.survey)
|
||||
|
||||
# Loop all the frequenies
|
||||
for freq in self.survey.freqs:
|
||||
dA_du = self.getA(freq) #
|
||||
dA_duI = self.Solver(dA_du, **self.solverOpts)
|
||||
|
||||
for src in self.survey.getSrcByFreq(freq):
|
||||
# We need fDeriv_m = df/du*du/dm + df/dm
|
||||
# Construct du/dm, it requires a solve
|
||||
ftype = self._fieldType + 'Solution'
|
||||
u_src = f[src, ftype]
|
||||
dA_dm = self.getADeriv_m(freq, u_src, v)
|
||||
@@ -44,28 +59,23 @@ class BaseMTProblem(BaseFDEMProblem):
|
||||
du_dm = dA_duI * ( - dA_dm )
|
||||
else:
|
||||
du_dm = dA_duI * ( - dA_dm + dRHS_dm )
|
||||
# Calculate the projection derivatives
|
||||
for rx in src.rxList:
|
||||
# df_duFun = u.deriv_u(rx.fieldsUsed, m)
|
||||
if 'e' in self._fieldType:
|
||||
projField = 'b'
|
||||
elif 'b' in self._fieldType:
|
||||
projField = 'e'
|
||||
df_duFun = getattr(f, '_%sDeriv_u'%projField, None)
|
||||
df_du = df_duFun(src, du_dm, adjoint=False)
|
||||
if df_du is not None:
|
||||
du_dm = df_du
|
||||
# Get the stacked derivative
|
||||
# df_duFun = getattr(f, '_fDeriv_u', None)
|
||||
# df_dmFun = getattr(f, '_fDeriv_m', None)
|
||||
# df_dm = df_dmFun(src,v,adjoint=False)
|
||||
# if df_dm is None:
|
||||
# fDeriv_m = df_duFun(src, du_dm, adjoint=False)
|
||||
# else:
|
||||
# fDeriv_m = df_duFun(src, du_dm, adjoint=False) + df_dm
|
||||
# Not needed for now. Since PDeriv does this currently.
|
||||
|
||||
df_dmFun = getattr(f, '_%sDeriv_m'%projField, None)
|
||||
df_dm = df_dmFun(src, v, adjoint=False)
|
||||
if df_dm is not None:
|
||||
du_dm += df_dm
|
||||
|
||||
P = lambda v: rx.projectFieldsDeriv(src, self.mesh, f, v) # wrt u, also have wrt m
|
||||
|
||||
|
||||
Jv[src, rx] = P(du_dm)
|
||||
|
||||
return Utils.mkvc(Jv)
|
||||
# Get the projection derivative
|
||||
PDeriv = lambda v: rx.projectFieldsDeriv(src, self.mesh, f, v) # wrt u, also have wrt m
|
||||
Jv[src, rx] = PDeriv(du_dm)
|
||||
# Return the vectorized sensitivities
|
||||
return mkvc(Jv)
|
||||
|
||||
def Jtvec(self, m, v, f=None):
|
||||
if f is None:
|
||||
@@ -88,29 +98,18 @@ class BaseMTProblem(BaseFDEMProblem):
|
||||
u_src = f[src, ftype]
|
||||
|
||||
for rx in src.rxList:
|
||||
# Get the adjoint projectFieldsDeriv
|
||||
PTv = rx.projectFieldsDeriv(src, self.mesh, f, v[src, rx], adjoint=True) # wrt u, need possibility wrt m
|
||||
|
||||
df_duTFun = getattr(f, '_%sDeriv_u'%rx.projField, None)
|
||||
df_duT = df_duTFun(src, PTv, adjoint=True)
|
||||
if df_duT is not None:
|
||||
dA_duIT = ATinv * df_duT
|
||||
else:
|
||||
dA_duIT = ATinv * PTv
|
||||
|
||||
# Get the
|
||||
dA_duIT = ATinv * PTv
|
||||
dA_dmT = self.getADeriv_m(freq, u_src, dA_duIT, adjoint=True)
|
||||
|
||||
dRHS_dmT = self.getRHSDeriv_m(src, dA_duIT, adjoint=True)
|
||||
|
||||
# Make du_dmT
|
||||
if dRHS_dmT is None:
|
||||
du_dmT = - dA_dmT
|
||||
else:
|
||||
du_dmT = -dA_dmT + dRHS_dmT
|
||||
|
||||
df_dmFun = getattr(f, '_%sDeriv_m'%rx.projField, None)
|
||||
dfT_dm = df_dmFun(src, PTv, adjoint=True)
|
||||
if dfT_dm is not None:
|
||||
du_dmT += dfT_dm
|
||||
|
||||
# Select the correct component
|
||||
real_or_imag = rx.projComp
|
||||
if real_or_imag == 'real':
|
||||
Jtv += du_dmT.real
|
||||
|
||||
+29
-4
@@ -80,10 +80,11 @@ class FieldsMT_1D(FieldsMT):
|
||||
return - 1./(1j*omega(src.freq)) * (C * v)
|
||||
|
||||
def _bSecondaryDeriv_m(self, src, v, adjoint = False):
|
||||
S_mDeriv, _ = src.evalDeriv(self.survey.prob, adjoint)
|
||||
S_mDeriv = S_mDeriv(v)
|
||||
if S_mDeriv is not None:
|
||||
return 1./(1j * omega(src.freq)) * S_mDeriv
|
||||
# Doesn't depend on m
|
||||
# _, S_eDeriv = src.evalDeriv(self.survey.prob, adjoint)
|
||||
# S_eDeriv = S_eDeriv(v)
|
||||
# if S_eDeriv is not None:
|
||||
# return 1./(1j * omega(src.freq)) * S_eDeriv
|
||||
return None
|
||||
|
||||
def _bDeriv_u(self, src, v, adjoint=False):
|
||||
@@ -94,3 +95,27 @@ class FieldsMT_1D(FieldsMT):
|
||||
# Assuming the primary does not depend on the model
|
||||
return self._bSecondaryDeriv_m(src, v, adjoint)
|
||||
|
||||
def _fDeriv_u(self, src, v, adjoint=False):
|
||||
"""
|
||||
Derivative of the fields object wrt u.
|
||||
|
||||
:param MTsrc src: MT source
|
||||
:param numpy.ndarray v: random vector of f_sol.size
|
||||
This function stacks the fields derivatives appropriately
|
||||
|
||||
return a vector of size (nreEle+nrbEle)
|
||||
"""
|
||||
|
||||
de_du = v #Utils.spdiag(np.ones((self.nF,)))
|
||||
db_du = self._bDeriv_u(src, v, adjoint)
|
||||
# Return the stack
|
||||
# This doesn't work...
|
||||
return np.vstack((de_du,db_du))
|
||||
|
||||
def _fDeriv_m(self, src, v, adjoint=False):
|
||||
"""
|
||||
Derivative of the fields object wrt m.
|
||||
|
||||
This function stacks the fields derivatives appropriately
|
||||
"""
|
||||
return None
|
||||
@@ -60,7 +60,8 @@ class eForm_psField(BaseMTProblem):
|
||||
if adjoint:
|
||||
return 1j * omega(freq) * ( dMf_dsig.T * v )
|
||||
# Note: output has to be nN/nF, not nC/nE.
|
||||
return 1j * omega(freq) * ( (dMf_dsig * dMf_dsig.T)**(1/2) * v)
|
||||
# v should be nC
|
||||
return 1j * omega(freq) * ( dMf_dsig * v )
|
||||
|
||||
def getRHS(self, freq):
|
||||
"""
|
||||
@@ -106,17 +107,11 @@ class eForm_psField(BaseMTProblem):
|
||||
|
||||
# Store the fields
|
||||
Src = self.survey.getSrcByFreq(freq)[0]
|
||||
# Calculate total e
|
||||
# NOTE: only store the e_solution(secondary), all other components calculated in the fields object
|
||||
F[Src, 'e_1dSolution'] = e_s[:,1] # Only storing the yx polarization as 1d
|
||||
|
||||
e = Src.ePrimary(self) + e_s
|
||||
|
||||
# Store the fields
|
||||
# NOTE: only store
|
||||
F[Src, 'e_1dSolution'] = e[:,1] # Only storing the yx polarization as 1d
|
||||
# F[Src, 'e_py'] = 0*e[:,0]
|
||||
# Note curl e = -iwb so b = -curl e /iw
|
||||
# b = -( self.mesh.nodalGrad * e )/( 1j*omega(freq) )
|
||||
# F[Src, 'b_px'] = 0*b[:,0]
|
||||
# F[Src, 'b_1d'] = b[:,1]
|
||||
if self.verbose:
|
||||
print 'Ran for {:f} seconds'.format(time.time()-startTime)
|
||||
|
||||
+28
-11
@@ -85,7 +85,7 @@ class RxMT(Survey.BaseRx):
|
||||
|
||||
def projectFields(self, src, mesh, f):
|
||||
'''
|
||||
Project the fields and return the
|
||||
Project the fields and return the correct data.
|
||||
'''
|
||||
|
||||
if self.projType is 'Z1D':
|
||||
@@ -131,32 +131,48 @@ class RxMT(Survey.BaseRx):
|
||||
def projectFieldsDeriv(self, src, mesh, f, v, adjoint=False):
|
||||
"""
|
||||
The derivative of the projection wrt u
|
||||
|
||||
:param MTsrc src: MT source
|
||||
:param TensorMesh mesh: Mesh defining the topology of the problem
|
||||
:param MTfields f: MT fields object of the source
|
||||
:param numpy.ndarray v: Random vector of size
|
||||
"""
|
||||
|
||||
real_or_imag = self.projComp
|
||||
|
||||
if not adjoint:
|
||||
if self.projType is 'Z1D':
|
||||
Pex = mesh.getInterpolationMat(self.locs,'Fx')
|
||||
Pbx = mesh.getInterpolationMat(self.locs,'Ex')
|
||||
# ex = Pex*mkvc(f[src,'e_1d'],2)
|
||||
# bx = Pbx*mkvc(f[src,'b_1d'],2)/mu_0
|
||||
deriv_complex = Utils.sdiag(1./(Pbx*mkvc(f[src,'b_1d'],2)/mu_0))*(Pex*v) - Utils.sdiag(Pex*mkvc(f[src,'e_1d'],2))*(Utils.sdiag(1./(Pbx*mkvc(f[src,'b_1d'],2)/mu_0)).T*Utils.sdiag(1./(Pbx*mkvc(f[src,'b_1d'],2)/mu_0)))*(Pbx*f._bDeriv_u(src,v)/mu_0)
|
||||
dP_de = Utils.sdiag(1./(Pbx*mkvc(f[src,'b_1d'],2)/mu_0))*(Pex*v)
|
||||
dP_db = - Utils.sdiag(Pex*mkvc(f[src,'e_1d'],2))*(Utils.sdiag(1./(Pbx*mkvc(f[src,'b_1d'],2)/mu_0)).T*Utils.sdiag(1./(Pbx*mkvc(f[src,'b_1d'],2)/mu_0)))*(Pbx*f._bDeriv_u(src,v)/mu_0)
|
||||
PDeriv_complex = np.sum((dP_de,dP_db))
|
||||
# elif self.projType is 'Z2D
|
||||
elif self.projType is 'Z3D':
|
||||
pass
|
||||
Pv = getattr(deriv_complex, real_or_imag)
|
||||
raise NotImplementedError('Has not be implement for full impedance tensor')
|
||||
Pv = np.array(getattr(PDeriv_complex, real_or_imag))
|
||||
elif adjoint:
|
||||
raise NotImplementedError('must be real or imag')
|
||||
|
||||
if self.projType is 'Z1D':
|
||||
Pex = mesh.getInterpolationMat(self.locs,'Fx')
|
||||
Pbx = mesh.getInterpolationMat(self.locs,'Ex')
|
||||
# ex = Pex*mkvc(f[src,'e_1d'],2)
|
||||
# bx = Pbx*mkvc(f[src,'b_1d'],2)/mu_0
|
||||
dP_deTv = mkvc(Pex.T*Utils.sdiag(1./(Pbx*mkvc(f[src,'b_1d'],2)/mu_0)).T*v,2)
|
||||
db_duv = Pbx.T/mu_0*Utils.sdiag(1./(Pbx*mkvc(f[src,'b_1d'],2)/mu_0))*(Utils.sdiag(1./(Pbx*mkvc(f[src,'b_1d'],2)/mu_0))).T*Utils.sdiag(Pex*mkvc(f[src,'e_1d'],2)).T
|
||||
dP_dbTv = -mkvc(f._bDeriv_u(src,db_duv,adjoint=True)*v,2)
|
||||
PDeriv_complex = np.sum((dP_deTv,dP_dbTv))
|
||||
elif self.projType is 'Z3D':
|
||||
raise NotImplementedError('must be real or imag')
|
||||
Pv = np.array(getattr(PDeriv_complex, real_or_imag))
|
||||
return Pv
|
||||
|
||||
|
||||
# Note: Might need to add tests to make sure that both polarization have the same rxList.
|
||||
|
||||
###############
|
||||
### Sources ###
|
||||
###############
|
||||
# Note: Should like inheret from FDEM
|
||||
|
||||
class srcMT(SrcFDEM): # Survey.BaseSrc):
|
||||
'''
|
||||
Sources for the MT problem.
|
||||
@@ -238,15 +254,16 @@ class srcMT_polxy_1Dprimary(srcMT):
|
||||
if problem.mesh.dim == 1:
|
||||
# Need to use the faceInnerProduct
|
||||
MsigmaDeriv = problem.mesh.getFaceInnerProductDeriv(problem.curModel.sigma)(self.ePrimary(problem)[:,-1]) * problem.curModel.sigmaDeriv
|
||||
MsigmaDeriv = ( MsigmaDeriv * MsigmaDeriv.T)**2
|
||||
# MsigmaDeriv = ( MsigmaDeriv * MsigmaDeriv.T)**2
|
||||
if problem.mesh.dim == 2:
|
||||
pass
|
||||
if problem.mesh.dim == 3:
|
||||
MsigmaDeriv = problem.MeSigmaDeriv(self.ePrimary(problem))
|
||||
if adjoint:
|
||||
#
|
||||
return MsigmaDeriv.T * v
|
||||
else:
|
||||
# Moved the v in front to make the multi work
|
||||
# v should be nC size
|
||||
return MsigmaDeriv * v
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user