(self, ws, return_feature=False, **block_kwargs)
| 500 | setattr(self, f'b{res}', block) |
| 501 | |
| 502 | def forward(self, ws, return_feature=False, **block_kwargs): |
| 503 | block_ws = [] |
| 504 | features = [] |
| 505 | with torch.autograd.profiler.record_function('split_ws'): |
| 506 | misc.assert_shape(ws, [None, self.num_ws, self.w_dim]) |
| 507 | ws = ws.to(torch.float32) |
| 508 | w_idx = 0 |
| 509 | for res in self.block_resolutions: |
| 510 | block = getattr(self, f'b{res}') |
| 511 | block_ws.append(ws.narrow(1, w_idx, block.num_conv + block.num_torgb)) |
| 512 | w_idx += block.num_conv |
| 513 | |
| 514 | x = img = None |
| 515 | for res, cur_ws in zip(self.block_resolutions, block_ws): |
| 516 | block = getattr(self, f'b{res}') |
| 517 | x, img = block(x, img, cur_ws, **block_kwargs) |
| 518 | features.append(x) |
| 519 | if return_feature: |
| 520 | return img, features |
| 521 | else: |
| 522 | return img |
| 523 | |
| 524 | def extra_repr(self): |
| 525 | return ' '.join([ |
nothing calls this directly
no outgoing calls
no test coverage detected