| 12 | |
| 13 | |
| 14 | class 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 |
no outgoing calls
no test coverage detected
searching dependent graphs…