From 04b250119fae834829e4b1de929a265f0615a15e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=95=D0=B2=D0=B3=D0=B5=D0=BD=D0=B8=D1=8F=20=D0=A1=D1=83?= =?UTF-8?q?=D1=85=D0=BE=D0=B4=D0=BE=D0=BB=D1=8C=D1=81=D0=BA=D0=B0=D1=8F?= Date: Thu, 23 Jan 2025 18:10:41 +0100 Subject: [PATCH] throw away similar stemmed words of lesser frequency --- mini_coil/data_pipeline/combine_models.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/mini_coil/data_pipeline/combine_models.py b/mini_coil/data_pipeline/combine_models.py index 079ad96..187f3d6 100644 --- a/mini_coil/data_pipeline/combine_models.py +++ b/mini_coil/data_pipeline/combine_models.py @@ -9,7 +9,7 @@ from mini_coil.data_pipeline.vocab_resolver import VocabResolver from mini_coil.model.encoder import Encoder from mini_coil.model.word_encoder import WordEncoder - +from py_rust_stemmers import SnowballStemmer def load_vocab(vocab_path): vocab = [] @@ -27,17 +27,25 @@ def main(): parser.add_argument("--output-path", type=str) parser.add_argument("--input-dim", type=int, default=512) parser.add_argument("--output-dim", type=int, default=4) + + stemmer = SnowballStemmer("english") args = parser.parse_args() vocab = load_vocab(args.vocab_path) filtered_vocab = [] + stemmed_words_vocab = set() for word in vocab: if word in english_stopwords: continue - model_path = os.path.join(args.models_dir, f"model-{word}.ptch") - if os.path.exists(model_path): - filtered_vocab.append(word) + stemmed_word = stemmer.stem_word(word) + if stemmed_word in stemmed_words_vocab: + continue + else: + stemmed_words_vocab.add(stemmed_word) + model_path = os.path.join(args.models_dir, f"model-{word}.ptch") + if os.path.exists(model_path): + filtered_vocab.append(word) params = [torch.zeros(args.input_dim, args.output_dim)] # Extra zero tensor, as first word is vocab starts from 1