-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathModelBassPca_func_mf.m
More file actions
160 lines (148 loc) · 5.03 KB
/
Copy pathModelBassPca_func_mf.m
File metadata and controls
160 lines (148 loc) · 5.03 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
classdef ModelBassPca_func_mf < handle
% PCA Based Model Emulator
properties
model
mod_corr
stochastic
nmcmc
input_names
basis
meas_error_cor
discrep_cov
ii
trunc_error_var
mod_s2
emu_vars
yobs
marg_lik_cov
discrep_vars
nd
discrep_tau
D
discrep
nexp
exp_ind
s2
end
methods
function obj = ModelBassPca_func_mf(bmod, bmod_corr, input_names, exp_ind, s2)
% **PCA Based Model Emulator using BASS MultiFidelity**
%
% This function setups up emulator object
%
% bmod: a object of the type BassBasis
% bmod_corr: a cell array of objects of the type BassBasis,
% these are the corrections to the LF bmod
% input_names: cell array of strings of input variable names
% exp_ind: experiment indices (default: NaN)
% s2: how to sample error variance (default: 'MH')
%
% returns an object of class ModelBassPca_func_mf
arguments
bmod BassBasis
bmod_corr cell
input_names
exp_ind = NaN;
s2 = 'MH';
end
obj.model = bmod;
obj.mod_corr = bmod_corr;
obj.stochastic = true;
obj.nmcmc = length(bmod.bm_list{1}.samples.s2);
obj.input_names = input_names;
obj.basis = obj.model.basis;
obj.meas_error_cor = eye(size(obj.basis,1));
obj.discrep_cov = eye(size(obj.basis,1))*1e-12;
obj.ii = 1;
npc = obj.model.nbasis;
if npc > 1
obj.trunc_error_var = diag(cov(obj.model.trunc_error'));
else
obj.trunc_error_var = diag(reshape(cov(obj.model.trunc_error'),1,1));
end
obj.mod_s2 = zeros(obj.nmcmc, npc);
for i = 1:npc
obj.mod_s2(:,i) = obj.model.bm_list{i}.samples.s2;
end
obj.emu_vars = obj.mod_s2(obj.ii,:);
obj.yobs = NaN;
obj.marg_lik_cov = NaN;
obj.discrep_vars = NaN;
obj.nd = 0;
obj.discrep_tau = 1.;
obj.D = NaN;
obj.discrep = 0.;
if isnan(exp_ind)
exp_ind = 1;
end
obj.nexp = max(exp_ind);
obj.exp_ind = exp_ind;
obj.s2 = s2;
if strcmp(s2,'gibbs')
error( "Cannot use Gibbs s2 for emulator models.")
end
end
function obj = step(obj)
obj.ii = randsample(1:obj.nmcmc,1);
obj.emu_vars = obj.mod_s2(obj.ii,:);
end
function discrep_vars = discrep_sample(obj, yobs, pred, cov, itemp)
S = eye(obj.nd) ./ obj.discrep_tau + obj.D'*cov.inv*obj.D;
m = obj.D' * cov.inv * (yobs-pred)';
discrep_vars = chol_sample(S\m, S./itemp);
end
function pred = eval(obj, parmat, pool, nugget)
arguments
obj
parmat
pool = true;
nugget = false;
end
fn = obj.input_names;
parmat_array = zeros(length(parmat.(fn{1})),numel(fn));
for i = 1:numel(fn)
parmat_array(:,i) = parmat.(fn{i});
end
if pool
pred = obj.model.predict(parmat_array, obj.ii, nugget);
for i = 1:length(obj.mod_corr)
pred1 = obj.mod_corr{i}.predict(parmat_array, obj.ii, nugget);
if mod(i,2) == 1
pred = pred + pred1;
else
gam = v_to_gam(pred1');
for j = 1:size(gam,2)
pred(j,:) = warp_f_gamma(pred(j,:),gam(:,j),linspace(0,1,size(gam,1)))';
end
end
end
else
keyboard
end
end
function out = llik(~, yobs, pred, cov)
vec = yobs(:) - pred(:);
out = -0.5*(cov.ldet + vec'*cov.inv*vec);
end
function out = lik_cov_inv(obj, s2vec)
vec = obj.trunc_error_var + s2vec(:);
Ainv = diag(1./vec);
Aldet = sum(log(vec));
out = obj.swm(Ainv, obj.basis, diag(1./obj.emu_vars), obj.basis', Aldet, sum(log(obj.emu_vars)));
end
function out = chol_solve(~, x)
R = chol(x);
ldet = 2 * sum(log(diag(R)));
inv1 = inv(x);
out.inv = inv1;
out.ldet = ldet;
end
function out = swm(obj, Ainv, U, Cinv, V, Aldet, Cldet)
in_mat = obj.chol_solve(Cinv + V*Ainv*U);
inv1 = Ainv - Ainv * U * in_mat.inv * V * Ainv;
ldet = in_mat.ldet + Aldet + Cldet;
out.inv = inv1;
out.ldet = ldet;
end
end
end