|
1 | | -function f = mtimes(f,g) |
2 | | -% * Matrix multiplication for kernel |
| 1 | +function out = mtimes(f, g) |
| 2 | +% * Matrix multiplication for kernel objects. |
3 | 3 | % |
4 | | -% Currently only supports scalars |
| 4 | +% Scalar c * K or K * c: scales eval, shifted_eval, fmm (same as times). |
5 | 5 | % |
6 | | -% returns c*F or F*C for the kernel F and scalar c |
7 | | - f = times(f,g); |
| 6 | +% Left multiply M * K: M is a (p x m) matrix or function handle M(t) |
| 7 | +% returning (p x m x nt). Output opdims = [p, K.opdims(2)]. |
| 8 | +% |
| 9 | +% Right multiply K * N: N is a (q x p) matrix or function handle N(s) |
| 10 | +% returning (q x p x ns), q = K.opdims(2). Output opdims = [K.opdims(1), p]. |
| 11 | + |
| 12 | +% Determine which argument is the kernel and which is the multiplier. |
| 13 | +if isa(f, 'kernel') && isa(g, 'kernel') |
| 14 | + error('KERNEL:mtimes:invalid', ... |
| 15 | + 'Cannot * two kernel objects; use + to combine.'); |
| 16 | +end |
| 17 | + |
| 18 | +if ~isa(f, 'kernel') |
| 19 | + [f, g] = deal(g, f); |
| 20 | + side = 'left'; |
| 21 | +elseif ~isa(g, 'kernel') |
| 22 | + side = 'right'; |
| 23 | +else |
| 24 | + error('KERNEL:mtimes:invalid', 'Unexpected argument types.'); |
| 25 | +end |
| 26 | + |
| 27 | +K = f; |
| 28 | +h = g; |
| 29 | + |
| 30 | +% scalar: delegate to times |
| 31 | +if isnumeric(h) && isscalar(h) |
| 32 | + out = times(K, h); |
| 33 | + return; |
| 34 | +end |
| 35 | + |
| 36 | +% constant matrix: |
| 37 | +if isnumeric(h) && ~isscalar(h) |
| 38 | + A = h; |
| 39 | + if strcmp(side, 'left') |
| 40 | + assert(size(A,2) == K.opdims(1), ... |
| 41 | + 'KERNEL:mtimes: left matrix must have %d columns', K.opdims(1)); |
| 42 | + p = size(A, 1); |
| 43 | + else |
| 44 | + assert(size(A,1) == K.opdims(2), ... |
| 45 | + 'KERNEL:mtimes: right matrix must have %d rows', K.opdims(2)); |
| 46 | + p = size(A, 2); |
| 47 | + end |
| 48 | + h = A; |
| 49 | +end |
| 50 | + |
| 51 | +% only remaining option is a function handle |
| 52 | +if ~isa(h, 'function_handle') && ~isnumeric(h) |
| 53 | + error('KERNEL:mtimes:invalid', ... |
| 54 | + 'Argument must be a scalar, matrix, or function handle.'); |
| 55 | +end |
| 56 | + |
| 57 | +if isa(h, 'function_handle') |
| 58 | + nargfunc = nargin(h); |
| 59 | + assert(nargfunc==1, ... |
| 60 | + 'KERNEL:mtimes h must be a function of source or target, not both') |
| 61 | +end |
| 62 | + |
| 63 | +% probe h to determine output dimension p |
| 64 | +% |
| 65 | +% A function handle returning a (1 x 1 x n) array is treated as a |
| 66 | +% pointwise scalar multiplier |
| 67 | +ispointwise = false; |
| 68 | +if ~exist('p', 'var') |
| 69 | + try |
| 70 | + probe.r = randn(2,1); probe.d = randn(2,1); |
| 71 | + probe.d2 = randn(2,1); probe.n = randn(2,1); |
| 72 | + hval = h(probe); |
| 73 | + if size(hval,1) == 1 && size(hval,2) == 1 |
| 74 | + ispointwise = true; |
| 75 | + if strcmp(side, 'left') |
| 76 | + p = K.opdims(1); |
| 77 | + else |
| 78 | + p = K.opdims(2); |
| 79 | + end |
| 80 | + elseif strcmp(side, 'left') |
| 81 | + p = size(hval, 1); |
| 82 | + else |
| 83 | + p = size(hval, 2); |
| 84 | + end |
| 85 | + catch |
| 86 | + error('KERNEL:mtimes:probe', ... |
| 87 | + 'Could not probe function handle to determine output dimension.'); |
| 88 | + end |
| 89 | +end |
| 90 | + |
| 91 | +if ~ispointwise && exist('hval', 'var') |
| 92 | + if strcmp(side, 'left') |
| 93 | + assert(size(hval,2) == K.opdims(1), ... |
| 94 | + 'KERNEL:mtimes: left function handle must return matrices with %d columns', K.opdims(1)); |
| 95 | + else |
| 96 | + assert(size(hval,1) == K.opdims(2), ... |
| 97 | + 'KERNEL:mtimes: right function handle must return matrices with %d rows', K.opdims(2)); |
| 98 | + end |
| 99 | +end |
| 100 | + |
| 101 | +Keval = K.eval; |
| 102 | +Kshifted_eval = K.shifted_eval; |
| 103 | +Kfmm = K.fmm; |
| 104 | +m = K.opdims(1); |
| 105 | +q = K.opdims(2); |
| 106 | + |
| 107 | +out = K; |
| 108 | +out.type = ['custom_', K.type]; |
| 109 | +out.name = ['custom ', K.name]; |
| 110 | + |
| 111 | + function fval = evalh(pts) |
| 112 | + % h is either a constant (p x m) or (q x p) matrix |
| 113 | + if isnumeric(h) |
| 114 | + fval = h; |
| 115 | + else |
| 116 | + fval = h(pts); |
| 117 | + end |
| 118 | + end |
| 119 | + |
| 120 | + function pts = shift_pts(pts, o) |
| 121 | + % Translate the positions of pts by o |
| 122 | + pts.r = pts.r + o(:); |
| 123 | + end |
| 124 | + |
| 125 | + function out = apply_left(fval, X) |
| 126 | + if size(fval,1) == 1 && size(fval,2) == 1 |
| 127 | + out = fval .* X; |
| 128 | + else |
| 129 | + out = pagemtimes(fval, X); |
| 130 | + end |
| 131 | + end |
| 132 | + |
| 133 | + function out = apply_right(X, fval) |
| 134 | + if size(fval,1) == 1 && size(fval,2) == 1 |
| 135 | + out = X .* fval; |
| 136 | + else |
| 137 | + out = pagemtimes(X, fval); |
| 138 | + end |
| 139 | + end |
| 140 | + |
| 141 | +if strcmp(side, 'left') |
| 142 | + out.opdims = [p, q]; |
| 143 | + out.eval = @eval_left; |
| 144 | + out.fmm = set_if_exist(Kfmm, @fmm_left); |
| 145 | + out.shifted_eval = set_if_exist(Kshifted_eval, @shifted_eval_left); |
| 146 | +else |
| 147 | + out.opdims = [m, p]; |
| 148 | + out.eval = @eval_right; |
| 149 | + out.fmm = set_if_exist(Kfmm, @fmm_right); |
| 150 | + out.shifted_eval = set_if_exist(Kshifted_eval, @shifted_eval_right); |
| 151 | +end |
| 152 | + |
| 153 | +% left-multiply: h(t) * K(s,t) |
| 154 | + |
| 155 | + function fval4 = reshape_fval_left(fval, nt) |
| 156 | + % fval is either p x m x nt (per-target) or p x m (constant); |
| 157 | + % reshape to p x m x nt x 1, broadcasting the constant case. |
| 158 | + if size(fval, 3) == nt |
| 159 | + fval4 = reshape(fval, size(fval,1), size(fval,2), nt, 1); |
| 160 | + else |
| 161 | + fval4 = reshape(fval, size(fval,1), size(fval,2), 1, 1); |
| 162 | + end |
| 163 | + end |
| 164 | + |
| 165 | + function vals = eval_left(s, t) |
| 166 | + nt = size(t.r, 2); |
| 167 | + fval = evalh(t); |
| 168 | + Kmat = Keval(s, t); |
| 169 | + K4 = reshape(Kmat, m, 1, nt, []); |
| 170 | + out4 = apply_left(reshape_fval_left(fval, nt), K4); |
| 171 | + vals = reshape(out4, p*nt, []); |
| 172 | + end |
| 173 | + |
| 174 | + function vals = shifted_eval_left(s, t, o) |
| 175 | + nt = size(t.r, 2); |
| 176 | + fval = evalh(shift_pts(t, o)); |
| 177 | + Kmat = Kshifted_eval(s, t, o); |
| 178 | + K4 = reshape(Kmat, m, 1, nt, []); |
| 179 | + out4 = apply_left(reshape_fval_left(fval, nt), K4); |
| 180 | + vals = reshape(out4, p*nt, []); |
| 181 | + end |
| 182 | + |
| 183 | + function out = fmm_left(eps, s, t, sigma) |
| 184 | + nt = size(t.r, 2); |
| 185 | + fval = evalh(t); |
| 186 | + inner = Kfmm(eps, s, t, sigma); |
| 187 | + out = reshape(apply_left(fval, reshape(inner, m, 1, nt)), p*nt, 1); |
| 188 | + end |
| 189 | + |
| 190 | +% right-multiply: K(s,t) * h(s) |
| 191 | + |
| 192 | + function vals = eval_right(s, t) |
| 193 | + ns = size(s.r, 2); |
| 194 | + nt = size(t.r, 2); |
| 195 | + fval = evalh(s); |
| 196 | + Kmat = Keval(s, t); |
| 197 | + K3 = reshape(Kmat, m*nt, q, ns); |
| 198 | + vals = reshape(apply_right(K3, fval), m*nt, p*ns); |
| 199 | + end |
| 200 | + |
| 201 | + function vals = shifted_eval_right(s, t, o) |
| 202 | + ns = size(s.r, 2); |
| 203 | + nt = size(t.r, 2); |
| 204 | + fval = evalh(shift_pts(s, o)); |
| 205 | + Kmat = Kshifted_eval(s, t, o); |
| 206 | + K3 = reshape(Kmat, m*nt, q, ns); |
| 207 | + vals = reshape(apply_right(K3, fval), m*nt, p*ns); |
| 208 | + end |
| 209 | + |
| 210 | + function out = fmm_right(eps, s, t, sigma) |
| 211 | + ns = size(s.r, 2); |
| 212 | + fval = evalh(s); |
| 213 | + sig_in = reshape(apply_left(fval, reshape(sigma, p, 1, ns)), q, ns); |
| 214 | + out = Kfmm(eps, s, t, sig_in); |
| 215 | + end |
| 216 | + |
| 217 | +end |
| 218 | + |
| 219 | +function out = set_if_exist(cond, val) |
| 220 | +if isa(cond, 'function_handle') |
| 221 | + out = val; |
| 222 | +else |
| 223 | + out = []; |
| 224 | +end |
8 | 225 | end |
0 commit comments