Files
2026-03-13 11:24:41 +08:00

325 lines
10 KiB
Matlab

classdef ADI
% ADI class: simple implementation of automatic differentiation for easy construction of jacobian matrices.
%
% SYNOPSIS:
% x = ADI(value, jacobian)
%
% PARAMETERS:
% value - The numerical value of the object
%
% jacobian - The Jacobian of the object.
%
% RETURNS:
% ADI object.
%
% COMMENTS:
% This class is typically instansiated for a set of different variables
% using initVariablesADI. The file contains a worked example demonstrating
% the usage for several variables.
%
% SEE ALSO:
% initVariablesADI
%{
Copyright 2009-2014 SINTEF ICT, Applied Mathematics.
This file is part of The MATLAB Reservoir Simulation Toolbox (MRST).
MRST is free software: you can redistribute it and/or modify
it under the terms of the GNU General Public License as published by
the Free Software Foundation, either version 3 of the License, or
(at your option) any later version.
MRST is distributed in the hope that it will be useful,
but WITHOUT ANY WARRANTY; without even the implied warranty of
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
GNU General Public License for more details.
You should have received a copy of the GNU General Public License
along with MRST. If not, see <http://www.gnu.org/licenses/>.
%}
properties
val %function value
jac %list of sparse jacobian matrices
end
methods
function obj = ADI(a,b)
%ADI class constructor
if nargin == 0 % empty constructor
obj.val = [];
obj.jac = {};
elseif nargin == 1 %
if isa(a, 'ADI')
obj = a;
else
error('Contructor requires 2 inputs')
end
elseif nargin == 2 % values + jacobians
obj.val = a; % value
if ~iscell(b)
b = {b};
end
obj.jac = b; % jacobian or list of jacobians
else
error('Input to constructor not valid')
end
end
%--------------------------------------------------------------------
function h = numval(u)
h = numel(u.val);
end
%--------------------------------------------------------------------
function h = double(u)
h = u.val;
end
%--------------------------------------------------------------------
function h = ge(u, v)
h = ge(double(u), double(v));
end
%--------------------------------------------------------------------
function h = gt(u, v)
h = gt(double(u), double(v));
end
%--------------------------------------------------------------------
function h = le(u, v)
h = le(double(u), double(v));
end
%--------------------------------------------------------------------
function h = lt(u, v)
h = lt(double(u), double(v));
end
%--------------------------------------------------------------------
function h = uplus(u)
h = u;
end
%--------------------------------------------------------------------
function h = uminus(u)
h = ADI(-u.val, uminusJac(u.jac));
end
%--------------------------------------------------------------------
% function h = plus(u,v)
% if ~isa(u,'ADI') %u is a vector/scalar
% h = ADI(u+v.val, v.jac);
% elseif ~isa(v,'ADI') %v is a vector/scalar
% h = ADI(u.val + v, u.jac);
% else
% h = ADI(u.val+v.val, plusJac(u.jac, v.jac) );
% end
% end
function h = plus(u,v)
if ~isa(u,'ADI') %u is a vector/scalar
if numel(u) <= numel(v.val)
h = ADI(u+v.val, v.jac);
elseif numel(v.val) == 1
h = plus(u, repmat(v,[numel(u), 1]));
else
error('Vectors have different lengths')
end
elseif ~isa(v,'ADI') %v is a vector/scalar
if numel(v) <= numel(u.val)
h = ADI(u.val + v, u.jac);
elseif numel(u.val) == 1
h = plus(repmat(u,[numel(v), 1]), v);
else
error('Vectors have different lengths')
end
else
if numel(u.val) == numel(v.val)
h = ADI(u.val+v.val, plusJac(u.jac, v.jac) );
elseif numel(u.val) == 1
h = plus(repmat(u, [numel(v.val), 1]), v);
elseif numel(v.val) == 1
h = plus(u, repmat(v, [numel(u.val), 1]));
else
error('Vectors have different lengths')
end
end
end
%--------------------------------------------------------------------
function h = minus(u,v)
h = plus(u, uminus(v));
end
%--------------------------------------------------------------------
function h = mtimes(u,v)% '*'
if ~isa(u,'ADI') %u is a scalar/matrix
h = ADI(u*v.val, mtimesJac(u, v.jac));
elseif ~isa(v,'ADI') %v is a scalar
h = mtimes(v,u);
else % special case where either u or v has single value
if numel(u.val) == 1
h = times(repmat(u, [numel(v.val), 1]), v);
elseif numel(v.val) == 1
h = times(u, repmat(v, [numel(u.val), 1]));
else
error('Operation not supported');
end
end
end
%--------------------------------------------------------------------
function h = times(u,v)% '.*'
if ~isa(u,'ADI') %u is a scalar/vector
if numel(u)==numel(v.val)
h = ADI(u.*v.val, lMultDiag(u, v.jac));
else
h = mtimes(u,v);
end
elseif ~isa(v,'ADI') %v is a scalar/vector
h = times(v,u);
else
if numel(u.val)==numel(v.val)
h = ADI(u.val.*v.val, timesJac(u.val, v.val, u.jac, v.jac));
elseif numel(v.val)==1||numel(u.val)==1
h = mtimes(u,v);
else
error('Operation not supported');
end
end
end
%--------------------------------------------------------------------
function h = mrdivide(u,v)% '/'
if ~isa(v,'ADI') %v is a scalar
h = mtimes(u, 1/v);
else
error('Operation not supported');
end
end
%--------------------------------------------------------------------
function h = subsref(u,s)
switch s(1).type
case '.'
h = builtin('subsref',u,s);
case '()'
assert(numel(s(1).subs) == 1, ...
'Expected single index, got %d', numel(s(1).subs))
subs = s(1).subs{1};
if ischar(s) && strcmp(subs, ':'),
h = u;
else
if islogical(subs), subs = find(subs); end
h = ADI(u.val(subs), subsrefJac(u.jac, subs));
end
if numel(s) > 1
% Recursively handle next operation
h = subsref(h, s(2:end));
end
case '{}'
error('Operation not supported');
end
end
%--------------------------------------------------------------------
function h = power(u,v)% '.^'
h = ADI(u.val.^v, lMultDiag(v.*u.val.^(v-1), u.jac));
end
%--------------------------------------------------------------------
function h = rdivide(u,v)% './'
h = times(u, power(v, -1));
end
%--------------------------------------------------------------------
function h = exp(u)
eu = exp(u.val);
h = ADI(eu, lMultDiag(eu, u.jac));
end
%--------------------------------------------------------------------
function h = log(u)
logu = log(u.val);
h = ADI(logu, lMultDiag(1./u.val, u.jac));
end
%--------------------------------------------------------------------
function h = interptable(X, Y, u)
y = interptable(X, Y, u.val);
dydx = dinterptable(X, Y, u.val);
h = ADI(y,lMultDiag(dydx, u.jac));
end
%--------------------------------------------------------------------
end
end
%**************************************************************************
%-------- Helper functions involving Jacobians ---------------------------
%**************************************************************************
function J = uminusJac(J1)
J = cellfun(@uminus, J1, 'UniformOutput', false);
end
function J = plusJac(J1, J2)
J = cellfun(@plus, J1, J2, 'UniformOutput', false);
end
function J = mtimesJac(M, J1)
J = cell(1, numel(J1));
for k = 1:numel(J)
J{k} = M*J1{k};
end
end
function J = lMultDiag(d, J1)
n = numel(d);
D = sparse((1:n)', (1:n)', d, n, n);
J = cell(1, numel(J1));
for k = 1:numel(J)
J{k} = D*J1{k};
end
end
function J = timesJac(v1, v2, J1, J2)
n = numel(v1);
D1 = sparse((1:n)', (1:n)', v1, n, n);
D2 = sparse((1:n)', (1:n)', v2, n, n);
J = cell(1, numel(J1));
for k = 1:numel(J)
J{k} = D1*J2{k} + D2*J1{k};
end
end
function J = subsrefJac(J1, subs)
J = cell(1, numel(J1));
for k = 1:numel(J)
J{k} = J1{k}(subs,:);
end
end
%--------------------------------------------------------------------------
%--------------------------------------------------------------------------
%--------------------------------------------------------------------------