Transforms for SAM2.
(
self, resolution, mask_threshold, max_hole_area=0.0, max_sprinkle_area=0.0
)
| 12 | |
| 13 | class SAM2Transforms(nn.Module): |
| 14 | def __init__( |
| 15 | self, resolution, mask_threshold, max_hole_area=0.0, max_sprinkle_area=0.0 |
| 16 | ): |
| 17 | """ |
| 18 | Transforms for SAM2. |
| 19 | """ |
| 20 | super().__init__() |
| 21 | self.resolution = resolution |
| 22 | self.mask_threshold = mask_threshold |
| 23 | self.max_hole_area = max_hole_area |
| 24 | self.max_sprinkle_area = max_sprinkle_area |
| 25 | self.mean = [0.485, 0.456, 0.406] |
| 26 | self.std = [0.229, 0.224, 0.225] |
| 27 | self.to_tensor = ToTensor() |
| 28 | self.transforms = torch.jit.script( |
| 29 | nn.Sequential( |
| 30 | Resize((self.resolution, self.resolution)), |
| 31 | Normalize(self.mean, self.std), |
| 32 | ) |
| 33 | ) |
| 34 | |
| 35 | def __call__(self, x): |
| 36 | x = self.to_tensor(x) |
nothing calls this directly
no outgoing calls
no test coverage detected