(
self,
params: Union[Dict[str, torch.nn.Parameter], torch.nn.ParameterDict],
optimizers: Dict[str, torch.optim.Optimizer],
state: Dict[str, Any],
maxbounds,
minbounds
)
| 377 | |
| 378 | @torch.no_grad() |
| 379 | def removeminmax( |
| 380 | self, |
| 381 | params: Union[Dict[str, torch.nn.Parameter], torch.nn.ParameterDict], |
| 382 | optimizers: Dict[str, torch.optim.Optimizer], |
| 383 | state: Dict[str, Any], |
| 384 | maxbounds, |
| 385 | minbounds |
| 386 | ): |
| 387 | maxx, maxy, maxz = maxbounds |
| 388 | minx, miny, minz = minbounds |
| 389 | xyz = params["means"] |
| 390 | mask0 = xyz[:,0] > maxx.item() |
| 391 | mask1 = xyz[:,1] > maxy.item() |
| 392 | mask2 = xyz[:,2] > maxz.item() |
| 393 | |
| 394 | mask3 = xyz[:,0] < minx.item() |
| 395 | mask4 = xyz[:,1] < miny.item() |
| 396 | mask5 = xyz[:,2] < minz.item() |
| 397 | mask = self.logicalorlist([mask0, mask1, mask2, mask3, mask4, mask5]) |
| 398 | remove(params=params, optimizers=optimizers, state=state, mask=mask) |
| 399 | return |
| 400 | |
| 401 | |
| 402 |
no test coverage detected