(module, **kwargs)
| 55 | |
| 56 | # for gradient checkpointing |
| 57 | def create_custom_forward(module, **kwargs): |
| 58 | def custom_forward(*inputs): |
| 59 | return module(*inputs, **kwargs) |
| 60 | return custom_forward |
| 61 | |
| 62 | def get_clones(module, N): |
| 63 | return nn.ModuleList([copy.deepcopy(module) for i in range(N)]) |