@@ -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+
5883void 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
0 commit comments