Method
__init__
(
self,
config,
layer=3,
in_channels=13,
kd_flag=True,
num_agent=5,
compress_level=0,
train_completion=True,
)
Source from the content-addressed store, hash-verified
| 294 | """ |
| 295 | |
| 296 | def __init__( |
| 297 | self, |
| 298 | config, |
| 299 | layer=3, |
| 300 | in_channels=13, |
| 301 | kd_flag=True, |
| 302 | num_agent=5, |
| 303 | compress_level=0, |
| 304 | train_completion=True, |
| 305 | ): |
| 306 | super().__init__() |
| 307 | self.config = config |
| 308 | self.kd_flag = kd_flag |
| 309 | self.in_channels = in_channels |
| 310 | self.num_agent = num_agent |
| 311 | self.train_completion = train_completion |
| 312 | self.stpn = STPN_KD(config.map_dims[2], compress_level, train_completion) |
| 313 | |
| 314 | def get_feature_maps_size(self, feature_maps: tuple): |
| 315 | size = list(feature_maps.shape) |
Tested by
no test coverage detected