From 187a302be238d1e6a757871f77e42117ff221c64 Mon Sep 17 00:00:00 2001
From: shixian.shi <shixian.shi@alibaba-inc.com>
Date: 星期二, 27 六月 2023 16:11:26 +0800
Subject: [PATCH] update clas finetune
---
funasr/datasets/large_datasets/dataset.py | 18 ++++++++++--------
funasr/datasets/large_datasets/utils/tokenize.py | 3 +++
2 files changed, 13 insertions(+), 8 deletions(-)
diff --git a/funasr/datasets/large_datasets/dataset.py b/funasr/datasets/large_datasets/dataset.py
index 5f2c2c6..1e9bb26 100644
--- a/funasr/datasets/large_datasets/dataset.py
+++ b/funasr/datasets/large_datasets/dataset.py
@@ -202,14 +202,7 @@
data_types = conf.get("data_types", "kaldi_ark,text")
pre_hwfile = conf.get("pre_hwlist", None)
- pre_prob = conf.get("pre_prob", 0) # unused yet
-
- hw_config = {"sample_rate": conf.get("sample_rate", 0.6),
- "double_rate": conf.get("double_rate", 0.1),
- "hotword_min_length": conf.get("hotword_min_length", 2),
- "hotword_max_length": conf.get("hotword_max_length", 8),
- "pre_prob": conf.get("pre_prob", 0.0)}
-
+ # pre_prob = conf.get("pre_prob", 0) # unused yet
if pre_hwfile is not None:
pre_hwlist = []
with open(pre_hwfile, 'r') as fin:
@@ -218,6 +211,15 @@
else:
pre_hwlist = None
+ hw_config = {"sample_rate": conf.get("sample_rate", 0.6),
+ "double_rate": conf.get("double_rate", 0.1),
+ "hotword_min_length": conf.get("hotword_min_length", 2),
+ "hotword_max_length": conf.get("hotword_max_length", 8),
+ "pre_prob": conf.get("pre_prob", 0.0),
+ "pre_hwlist": pre_hwlist}
+
+
+
dataset = AudioDataset(scp_lists,
data_names,
data_types,
diff --git a/funasr/datasets/large_datasets/utils/tokenize.py b/funasr/datasets/large_datasets/utils/tokenize.py
index a7eb6d2..3128a06 100644
--- a/funasr/datasets/large_datasets/utils/tokenize.py
+++ b/funasr/datasets/large_datasets/utils/tokenize.py
@@ -54,6 +54,9 @@
length = len(text)
if 'hw_tag' in data:
+ if hw_config['pre_hwlist'] is not None and hw_config['pre_prob'] > 0:
+ # enable preset hotword detect in sampling
+ import pdb; pdb.set_trace()
hotword_indxs = sample_hotword(length, **hw_config)
data['hotword_indxs'] = hotword_indxs
del data['hw_tag']
--
Gitblit v1.9.1