MCPcopy Create free account
hub / github.com/rtqichen/torchdiffeq / BouncingBallExample

Class BouncingBallExample

examples/bouncing_ball.py:14–100  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

12
13
14class BouncingBallExample(nn.Module):
15 def __init__(self, radius=0.2, gravity=9.8, adjoint=False):
16 super().__init__()
17 self.gravity = nn.Parameter(torch.as_tensor([gravity]))
18 self.log_radius = nn.Parameter(torch.log(torch.as_tensor([radius])))
19 self.t0 = nn.Parameter(torch.tensor([0.0]))
20 self.init_pos = nn.Parameter(torch.tensor([10.0]))
21 self.init_vel = nn.Parameter(torch.tensor([0.0]))
22 self.absorption = nn.Parameter(torch.tensor([0.2]))
23 self.odeint = odeint_adjoint if adjoint else odeint
24
25 def forward(self, t, state):
26 pos, vel, log_radius = state
27 dpos = vel
28 dvel = -self.gravity
29 return dpos, dvel, torch.zeros_like(log_radius)
30
31 def event_fn(self, t, state):
32 # positive if ball in mid-air, negative if ball within ground.
33 pos, _, log_radius = state
34 return pos - torch.exp(log_radius)
35
36 def get_initial_state(self):
37 state = (self.init_pos, self.init_vel, self.log_radius)
38 return self.t0, state
39
40 def state_update(self, state):
41 """Updates state based on an event (collision)."""
42 pos, vel, log_radius = state
43 pos = (
44 pos + 1e-7
45 ) # need to add a small eps so as not to trigger the event function immediately.
46 vel = -vel * (1 - self.absorption)
47 return (pos, vel, log_radius)
48
49 def get_collision_times(self, nbounces=1):
50
51 event_times = []
52
53 t0, state = self.get_initial_state()
54
55 for i in range(nbounces):
56 event_t, solution = odeint_event(
57 self,
58 state,
59 t0,
60 event_fn=self.event_fn,
61 reverse_time=False,
62 atol=1e-8,
63 rtol=1e-8,
64 odeint_interface=self.odeint,
65 )
66 event_times.append(event_t)
67
68 state = self.state_update(tuple(s[-1] for s in solution))
69 t0 = event_t
70
71 return event_times

Callers 3

learn_physics.pyFile · 0.90
gradcheckFunction · 0.85
bouncing_ball.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…