Skip to content

Commit 3eb2a05

Browse files
committed
Add acoustic2d/3d CUDA wavelet gradients.
1 parent 73ab700 commit 3eb2a05

9 files changed

Lines changed: 667 additions & 31 deletions

File tree

src/sweep/csrc/common/cudautils.h

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,21 @@
22

33
#include <cuda_runtime.h>
44
#include <ATen/cuda/CUDAContext.h>
5+
#include <torch/extension.h>
6+
7+
#define SWEEP_CUDA_SYNC_CHECK(label) \
8+
do { \
9+
cudaError_t _launch_err = cudaGetLastError(); \
10+
TORCH_CHECK( \
11+
_launch_err == cudaSuccess, \
12+
label, " launch failed: ", cudaGetErrorString(_launch_err) \
13+
); \
14+
cudaError_t _sync_err = cudaStreamSynchronize(at::cuda::getCurrentCUDAStream()); \
15+
TORCH_CHECK( \
16+
_sync_err == cudaSuccess, \
17+
label, " execution failed: ", cudaGetErrorString(_sync_err) \
18+
); \
19+
} while (0)
520

621
struct AsyncCopyContext {
722
cudaStream_t compute_stream;

src/sweep/csrc/equations/acoustic2d/backward.cu

Lines changed: 132 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -55,9 +55,35 @@ void accumulate_imaging_2d(
5555
);
5656
}
5757

58+
void accumulate_source_gradient_2d(
59+
dim3 source_grid,
60+
dim3 source_block,
61+
const float* adjoint_ptr,
62+
const BackwardInput& p,
63+
torch::Tensor* grad_wavelet,
64+
int it,
65+
const SolverContext& ctx,
66+
int nsrc
67+
)
68+
{
69+
if (grad_wavelet == nullptr) {
70+
return;
71+
}
72+
73+
accumulate_source_grad_2d<<<source_grid, source_block>>>(
74+
adjoint_ptr,
75+
grad_wavelet->data_ptr<float>(),
76+
p.forward_sources_loc.data_ptr<int>(),
77+
it,
78+
nsrc,
79+
ctx
80+
);
81+
}
82+
5883
void run_full_imaging(
5984
const BackwardInput& p,
6085
torch::Tensor* grad,
86+
torch::Tensor* grad_wavelet,
6187
RTMOutput* rtm_out
6288
)
6389
{
@@ -70,7 +96,8 @@ void run_full_imaging(
7096
int C = vp.size(1);
7197
int nz = vp.size(2);
7298
int nx = vp.size(3);
73-
int nsrc = p.adjoint_source.size(1);
99+
int adjoint_nsrc = p.adjoint_sources_loc.size(1);
100+
int forward_nsrc = p.forward_sources_loc.size(1);
74101
int B = N * C;
75102

76103
int M = p.M;
@@ -89,7 +116,8 @@ void run_full_imaging(
89116
auto cpml = cpml_tensor.view();
90117

91118
auto launch_config = fdtd::Wave2D::make(nx, nz, B);
92-
auto source_config = fdtd::Geom::make(nsrc, B);
119+
auto adj_source_config = fdtd::Geom::make(adjoint_nsrc, B);
120+
auto forward_source_config = fdtd::Geom::make(forward_nsrc, B);
93121

94122
const int order =
95123
(M <= 4) ? static_cast<int>(2 * M) : -1;
@@ -121,17 +149,28 @@ void run_full_imaging(
121149
ctx
122150
);
123151

124-
add_source<<<source_config.grid, source_config.block>>>(
152+
add_source<<<adj_source_config.grid, adj_source_config.block>>>(
125153
adj_view.u_next,
126154
p.adjoint_source.data_ptr<float>(),
127155
p.adjoint_sources_loc.data_ptr<int>(),
128156
it,
129-
nsrc,
157+
adjoint_nsrc,
130158
ctx
131159
);
132160

133161
adjoint.swap();
134162

163+
accumulate_source_gradient_2d(
164+
forward_source_config.grid,
165+
forward_source_config.block,
166+
adjoint.u_now_t.data_ptr<float>(),
167+
p,
168+
grad_wavelet,
169+
it,
170+
ctx,
171+
forward_nsrc
172+
);
173+
135174
accumulate_imaging_2d(
136175
launch_config.grid,
137176
launch_config.block,
@@ -152,8 +191,9 @@ BackwardOutput backward(const BackwardInput& in)
152191
{
153192
BackwardOutput out;
154193
auto grad = torch::zeros_like(in.models[0]);
155-
run_full_imaging(in, &grad, nullptr);
156-
out.grads = {grad};
194+
auto grad_wavelet = torch::zeros_like(in.forward_source);
195+
run_full_imaging(in, &grad, &grad_wavelet, nullptr);
196+
out.grads = {grad_wavelet, grad};
157197
return out;
158198
}
159199

@@ -174,7 +214,7 @@ RTMOutput rtm(const BackwardInput& in)
174214

175215
RTMOutput out;
176216
init_rtm_output_2d(out, in.models[0]);
177-
run_full_imaging(in, nullptr, &out);
217+
run_full_imaging(in, nullptr, nullptr, &out);
178218
return out;
179219
}
180220

@@ -219,6 +259,7 @@ BackwardOutput backward_bs(const BackwardInput& in)
219259
forward.u_now_t.copy_(p.u_last_two.select(1,0).squeeze(0));
220260

221261
auto grad = torch::zeros_like(vp);
262+
auto grad_wavelet = torch::zeros_like(p.forward_source);
222263

223264
// For checking wavefields
224265
// torch::Tensor u_allt = torch::zeros({nt, B, 1, nz, nx}, vp.options());
@@ -301,6 +342,17 @@ BackwardOutput backward_bs(const BackwardInput& in)
301342
);
302343

303344
adjoint.swap();
345+
346+
accumulate_source_gradient_2d(
347+
fwd_source_config.grid,
348+
fwd_source_config.block,
349+
adjoint.u_now_t.data_ptr<float>(),
350+
p,
351+
&grad_wavelet,
352+
it,
353+
ctx,
354+
forward_nsrc
355+
);
304356

305357
ACOUSTIC2D_NOPML(
306358
order,
@@ -371,7 +423,49 @@ BackwardOutput backward_bs(const BackwardInput& in)
371423

372424
}
373425

374-
out.grads = {grad};
426+
if (p.nt > 0) {
427+
auto adj_view = adjoint.view();
428+
429+
ACOUSTIC2D(
430+
order,
431+
launch_config.grid,
432+
launch_config.block,
433+
adj_view,
434+
false,
435+
nullptr,
436+
vp.data_ptr<float>(),
437+
lap_ctx,
438+
grad_ctx,
439+
grad_ctx_x,
440+
grad_ctx_z,
441+
cpml,
442+
ctx
443+
);
444+
445+
add_source<<<adj_source_config.grid, adj_source_config.block>>>(
446+
adj_view.u_next,
447+
p.adjoint_source.data_ptr<float>(),
448+
p.adjoint_sources_loc.data_ptr<int>(),
449+
0,
450+
adjoint_nsrc,
451+
ctx
452+
);
453+
454+
adjoint.swap();
455+
456+
accumulate_source_gradient_2d(
457+
fwd_source_config.grid,
458+
fwd_source_config.block,
459+
adjoint.u_now_t.data_ptr<float>(),
460+
p,
461+
&grad_wavelet,
462+
0,
463+
ctx,
464+
forward_nsrc
465+
);
466+
}
467+
468+
out.grads = {grad_wavelet, grad};
375469
return out;
376470

377471
}
@@ -475,6 +569,7 @@ void process_recursive_interval_2d(
475569
const BackwardInput& p,
476570
const torch::Tensor& vp,
477571
torch::Tensor& grad,
572+
torch::Tensor& grad_wavelet,
478573
int order,
479574
dim3 wave_grid,
480575
dim3 wave_block,
@@ -561,6 +656,17 @@ void process_recursive_interval_2d(
561656

562657
adjoint.swap();
563658

659+
accumulate_source_gradient_2d(
660+
forward_source_grid,
661+
forward_source_block,
662+
adjoint.u_now_t.data_ptr<float>(),
663+
p,
664+
&grad_wavelet,
665+
start,
666+
ctx,
667+
forward_nsrc
668+
);
669+
564670
calculate_grad<<<wave_grid, wave_block>>>(
565671
u_this.data_ptr<float>(),
566672
adjoint.u_now_t.data_ptr<float>(),
@@ -604,6 +710,7 @@ void process_recursive_interval_2d(
604710
p,
605711
vp,
606712
grad,
713+
grad_wavelet,
607714
order,
608715
wave_grid,
609716
wave_block,
@@ -631,6 +738,7 @@ void process_recursive_interval_2d(
631738
p,
632739
vp,
633740
grad,
741+
grad_wavelet,
634742
order,
635743
wave_grid,
636744
wave_block,
@@ -699,6 +807,7 @@ BackwardOutput backward_ckpt(const BackwardInput& in)
699807
forward.allocate(vp, 2, true);
700808

701809
auto grad = torch::zeros_like(vp);
810+
auto grad_wavelet = torch::zeros_like(p.forward_source);
702811

703812
AcousticCPMLTensor cpml_tensor;
704813
cpml_tensor.allocate(p.pml_vals, 2);
@@ -790,6 +899,17 @@ BackwardOutput backward_ckpt(const BackwardInput& in)
790899

791900
adjoint.swap();
792901

902+
accumulate_source_gradient_2d(
903+
fwd_source_config.grid,
904+
fwd_source_config.block,
905+
adjoint.u_now_t.data_ptr<float>(),
906+
p,
907+
&grad_wavelet,
908+
it,
909+
ctx,
910+
forward_nsrc
911+
);
912+
793913
calculate_grad<<<launch_config.grid, launch_config.block>>>(
794914
chunk_forward[it - start].data_ptr<float>(),
795915
adjoint.u_now_t.data_ptr<float>(),
@@ -800,7 +920,7 @@ BackwardOutput backward_ckpt(const BackwardInput& in)
800920
}
801921
}
802922

803-
out.grads = {grad};
923+
out.grads = {grad_wavelet, grad};
804924
return out;
805925
}
806926

@@ -846,6 +966,7 @@ BackwardOutput backward_recursive_ckpt(const BackwardInput& in)
846966
zero_wavefield_state_2d(adjoint);
847967

848968
auto grad = torch::zeros_like(vp);
969+
auto grad_wavelet = torch::zeros_like(p.forward_source);
849970

850971
AcousticCPMLTensor cpml_tensor;
851972
cpml_tensor.allocate(p.pml_vals, 2);
@@ -892,6 +1013,7 @@ BackwardOutput backward_recursive_ckpt(const BackwardInput& in)
8921013
p,
8931014
vp,
8941015
grad,
1016+
grad_wavelet,
8951017
order,
8961018
launch_config.grid,
8971019
launch_config.block,
@@ -912,7 +1034,7 @@ BackwardOutput backward_recursive_ckpt(const BackwardInput& in)
9121034
);
9131035
}
9141036

915-
out.grads = {grad};
1037+
out.grads = {grad_wavelet, grad};
9161038
return out;
9171039
}
9181040

src/sweep/csrc/equations/acoustic2d/kernels.cu

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -96,3 +96,33 @@ __global__ void accumulate_rtm_image_2d(
9696
src_b[idx] += uf * uf;
9797
rec_b[idx] += ub * ub;
9898
}
99+
100+
__global__ void accumulate_source_grad_2d(
101+
const float* __restrict__ u_backward,
102+
float* __restrict__ grad_source,
103+
const int* __restrict__ sources_loc,
104+
int it,
105+
int nsrc,
106+
SolverContext solver
107+
) {
108+
int b = blockIdx.x;
109+
int s = blockIdx.y * blockDim.x + threadIdx.x;
110+
111+
if (b >= solver.B || s >= nsrc) {
112+
return;
113+
}
114+
115+
int base = (b * nsrc + s) * 2;
116+
int ix = sources_loc[base + 0];
117+
int iz = sources_loc[base + 1];
118+
119+
if (ix < 0 || ix >= solver.nx || iz < 0 || iz >= solver.nz) {
120+
return;
121+
}
122+
123+
int spatial_size = solver.nx * solver.nz;
124+
int u_idx = b * spatial_size + iz * solver.nx + ix;
125+
int grad_idx = (b * nsrc + s) * solver.nt + it;
126+
127+
grad_source[grad_idx] += u_backward[u_idx];
128+
}

src/sweep/csrc/equations/acoustic2d/kernels.cuh

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -182,3 +182,12 @@ __global__ void accumulate_rtm_image_2d(
182182
float* __restrict__ receiver_illumination,
183183
int nx, int nz
184184
);
185+
186+
__global__ void accumulate_source_grad_2d(
187+
const float* __restrict__ u_backward, // (B, nz, nx)
188+
float* __restrict__ grad_source, // (B, nsrc, nt)
189+
const int* __restrict__ sources_loc, // (B, nsrc, 2)
190+
int it,
191+
int nsrc,
192+
SolverContext solver
193+
);

0 commit comments

Comments
 (0)