Skip to content

Commit 5bae6b0

Browse files
committed
Add function handles to kernel mtimes
1 parent 02b3f38 commit 5bae6b0

3 files changed

Lines changed: 542 additions & 22 deletions

File tree

chunkie/@kernel/mtimes.m

Lines changed: 222 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,225 @@
1-
function f = mtimes(f,g)
2-
% * Matrix multiplication for kernel
1+
function out = mtimes(f, g)
2+
% * Matrix multiplication for kernel objects.
33
%
4-
% Currently only supports scalars
4+
% Scalar c * K or K * c: scales eval, shifted_eval, fmm (same as times).
55
%
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
8225
end

chunkie/demo/demo_clamped_plate_scatter.m

Lines changed: 20 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -37,18 +37,19 @@
3737

3838
% assembling system matrix
3939

40-
fkern = @(s,t) chnk.flex2d.kern(zk, s, t, 'clamped_plate');
41-
42-
kappa = signed_curvature(chnkr);
43-
kappa = kappa(:);
40+
fkern = kernel(@(s,t) chnk.flex2d.kern(zk, s, t, 'clamped_plate'));
4441

4542
opts = [];
4643
opts.sing = 'log';
4744

45+
% Left-multiply by M(t) = [-2, 0; -4*kappa(t), -2] so that the diagonal
46+
% term becomes the identity, i.e. sys = M*K + I.
47+
Mfun = @(t) reshape([-2*ones(1,size(t.r(:,:),2)); -4*chnk.curvature2d(t); ...
48+
zeros(1,size(t.r(:,:),2)); -2*ones(1,size(t.r(:,:),2))], 2, 2, []);
49+
4850
start = tic;
49-
sys = chunkermat(chnkr,fkern, opts);
50-
sys = sys - 0.5*eye(2*chnkr.npt);
51-
sys(2:2:end,1:2:end) = sys(2:2:end,1:2:end) + kappa.*eye(chnkr.npt);
51+
sys = chunkermat(chnkr, Mfun*fkern, opts);
52+
sys = sys + eye(2*chnkr.npt);
5253

5354
t1 = toc(start);
5455
fprintf('%5.2e s : time to assemble matrix\n',t1)
@@ -57,19 +58,25 @@
5758

5859
[r1, grad] = planewave(kvec, chnkr.r);
5960

60-
nx = chnkr.n(1,:);
61+
nx = chnkr.n(1,:);
6162
ny = chnkr.n(2,:);
6263

63-
normalderiv = grad(:, 1).*(nx.')+ grad(:, 2).*(ny.'); % Dirichlet and Neumann BC(Clamped BC)
64+
normalderiv = grad(:, 1).*(nx.')+ grad(:, 2).*(ny.'); % Dirichlet and Neumann BC(Clamped BC)
6465

6566
firstbc = -r1;
6667
secondbc = -normalderiv;
6768

6869
rhs = zeros(2*chnkr.npt, 1); rhs(1:2:end) = firstbc ; rhs(2:2:end) = secondbc;
6970

71+
% apply M(t) to the right-hand side
72+
ptinfo = []; ptinfo.r = chnkr.r(:,:); ptinfo.d = chnkr.d(:,:); ptinfo.d2 = chnkr.d2(:,:);
73+
kappa = chnk.curvature2d(ptinfo); kappa = kappa(:);
74+
rhs(2:2:end) = -4*kappa.*rhs(1:2:end) - 2*rhs(2:2:end);
75+
rhs(1:2:end) = -2*rhs(1:2:end);
76+
7077
% Solving linear system
7178

72-
start = tic; sol = gmres(sys,rhs,[],1e-12,100); t1 = toc(start);
79+
start = tic; sol = gmres(sys,rhs,[],1e-12,200); t1 = toc(start);
7380
fprintf('%5.2e s : time for dense gmres\n',t1)
7481

7582
% evaluate at targets and plot
@@ -109,7 +116,7 @@
109116
nexttile
110117
zztarg = nan(size(xxtarg));
111118
zztarg(out) = uin;
112-
h=pcolor(xxtarg,yytarg,imag(zztarg),"FaceColor","interp");
119+
h=pcolor(xxtarg,yytarg,imag(zztarg)); h.FaceColor="interp";
113120
set(h,'EdgeColor','none')
114121
clim([-maxu,maxu])
115122
colormap(redblue);
@@ -122,7 +129,7 @@
122129
nexttile
123130
zztarg = nan(size(xxtarg));
124131
zztarg(out) = uscat;
125-
h=pcolor(xxtarg,yytarg,imag(zztarg),"FaceColor","interp");
132+
h=pcolor(xxtarg,yytarg,imag(zztarg)); h.FaceColor="interp";
126133
set(h,'EdgeColor','none')
127134
clim([-maxu,maxu])
128135
colormap(redblue);
@@ -135,7 +142,7 @@
135142
nexttile
136143
zztarg = nan(size(xxtarg));
137144
zztarg(out) = utot;
138-
h=pcolor(xxtarg,yytarg,imag(zztarg),"FaceColor","interp");
145+
h=pcolor(xxtarg,yytarg,imag(zztarg)); h.FaceColor="interp";
139146
set(h,'EdgeColor','none')
140147
clim([-maxu,maxu])
141148
colormap(redblue);

0 commit comments

Comments
 (0)