From f8d1c79fe355efb18ae49e4363307dfec3ab89ce Mon Sep 17 00:00:00 2001
From: 雾聪 <wucong.lyb@alibaba-inc.com>
Date: 星期一, 07 八月 2023 16:14:11 +0800
Subject: [PATCH] Merge branch 'main' of https://github.com/alibaba-damo-academy/FunASR into main
---
egs/callhome/eend_ola/local/model_averaging.py | 28 ++++++++++++++++++++++++++++
1 files changed, 28 insertions(+), 0 deletions(-)
diff --git a/egs/callhome/eend_ola/local/model_averaging.py b/egs/callhome/eend_ola/local/model_averaging.py
new file mode 100755
index 0000000..1871cd9
--- /dev/null
+++ b/egs/callhome/eend_ola/local/model_averaging.py
@@ -0,0 +1,28 @@
+#!/usr/bin/env python3
+
+import argparse
+
+import torch
+
+
+def average_model(input_files, output_file):
+ output_model = {}
+ for ckpt_path in input_files:
+ model_params = torch.load(ckpt_path, map_location="cpu")
+ for key, value in model_params.items():
+ if key not in output_model:
+ output_model[key] = value
+ else:
+ output_model[key] += value
+ for key in output_model.keys():
+ output_model[key] /= len(input_files)
+ torch.save(output_model, output_file)
+
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ parser.add_argument("output_file")
+ parser.add_argument("input_files", nargs='+')
+ args = parser.parse_args()
+
+ average_model(args.input_files, args.output_file)
\ No newline at end of file
--
Gitblit v1.9.1