@@ -309,30 +309,74 @@ def do_nothing(state: deepwave.common.CallbackState) -> None:
309309 for i in range (len (out1 )):
310310 assert torch .allclose (out1 [i ], out2 [i ])
311311 assert torch .allclose (grad1_v , grad2_v )
312- assert torch .allclose (grad1_scatter , grad2_scatter )
313- v .grad .zero_ ()
314- scatter .grad .zero_ ()
312+
313+
314+ def test_scalar_born_backward_callback_only_call_count () -> None :
315+ """Check that the backward callback is called the correct number of times
316+ when no forward callback is provided.
317+ """
318+ v = torch .ones (10 , 10 ) * 1500
319+ scatter = torch .ones (10 , 10 )
320+ v .requires_grad_ ()
321+ scatter .requires_grad_ ()
322+ dx = 5.0
323+ dt = 0.004
324+ nt = 20
325+ source_amplitudes = torch .zeros (1 , 1 , nt )
326+ source_amplitudes [0 , 0 , 5 ] = 1
327+ source_locations = torch .zeros (1 , 1 , 2 , dtype = torch .long )
328+ source_locations [0 , 0 , 0 ] = 5
329+ source_locations [0 , 0 , 1 ] = 5
330+ receiver_locations = torch .zeros (1 , 1 , 2 , dtype = torch .long )
331+ receiver_locations [0 , 0 , 0 ] = 5
332+ receiver_locations [0 , 0 , 1 ] = 5
333+
334+ class Counter :
335+ """A simple counter class for callbacks."""
336+
337+ def __init__ (self ) -> None :
338+ self .count = 0
339+
340+ def __call__ (self , state : deepwave .common .CallbackState ) -> None :
341+ """Increments the counter."""
342+ self .count += 1
343+
344+ # Test with a frequency that divides nt evenly
345+ backward_counter = Counter ()
346+ out = deepwave .scalar_born (
347+ v ,
348+ scatter ,
349+ dx ,
350+ dt ,
351+ source_amplitudes = source_amplitudes ,
352+ source_locations = source_locations ,
353+ receiver_locations = receiver_locations ,
354+ backward_callback = backward_counter ,
355+ callback_frequency = 2 ,
356+ )
357+ out [- 1 ].sum ().backward ()
358+ grad1_scatter = scatter .grad .detach ()
359+ assert backward_counter .count == nt / 2
315360
316361 # Test with a frequency that does not divide nt evenly
317- out3 = deepwave .scalar_born (
362+ v .grad .zero_ ()
363+ scatter .grad .zero_ ()
364+ backward_counter = Counter ()
365+ out = deepwave .scalar_born (
318366 v ,
319367 scatter ,
320368 dx ,
321369 dt ,
322370 source_amplitudes = source_amplitudes ,
323371 source_locations = source_locations ,
324372 receiver_locations = receiver_locations ,
325- forward_callback = do_nothing ,
326- backward_callback = do_nothing ,
373+ backward_callback = backward_counter ,
327374 callback_frequency = 3 ,
328375 )
329- out3 [- 1 ].sum ().backward ()
330- grad3_v = v .grad .clone ()
331- grad3_scatter = scatter .grad .clone ()
332- for i in range (len (out1 )):
333- assert torch .allclose (out1 [i ], out3 [i ])
334- assert torch .allclose (grad1_v , grad3_v )
335- assert torch .allclose (grad1_scatter , grad3_scatter )
376+ out [- 1 ].sum ().backward ()
377+ grad2_scatter = scatter .grad .detach ()
378+ assert backward_counter .count == (nt + 2 ) // 3
379+ assert torch .allclose (grad1_scatter , grad2_scatter )
336380
337381
338382def test_scalar_born_multishot_equivalence () -> None :
0 commit comments