| 9 | |
| 10 | |
| 11 | def export_rebuild_model(model, **kwargs): |
| 12 | model.device = kwargs.get("device") |
| 13 | model.make_pad_mask = sequence_mask(kwargs["max_seq_len"], flip=False) |
| 14 | model.forward = types.MethodType(export_forward, model) |
| 15 | model.export_dummy_inputs = types.MethodType(export_dummy_inputs, model) |
| 16 | model.export_input_names = types.MethodType(export_input_names, model) |
| 17 | model.export_output_names = types.MethodType(export_output_names, model) |
| 18 | model.export_dynamic_axes = types.MethodType(export_dynamic_axes, model) |
| 19 | model.export_name = types.MethodType(export_name, model) |
| 20 | return model |
| 21 | |
| 22 | def export_forward( |
| 23 | self, |