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