(path, weight)
| 279 | return x |
| 280 | |
| 281 | def sharding(path, weight): |
| 282 | parts = path.split(".") |
| 283 | even = int(parts[1]) % 2 == 0 |
| 284 | if even: |
| 285 | return 0 |
| 286 | else: |
| 287 | return -1 if parts[-1] != "bias" else None |
| 288 | |
| 289 | mod = nn.Sequential( |
| 290 | MyConv(3, 128, kernel_size=3), |