Word2Vec单词分类脚本:输入读取与结果存储问题求助
Word2Vec单词分类脚本修改问题
原代码及运行输出
原代码
# Category -> words data = { 'Names': ['john','jay','dan','nathan','bob'], 'Colors': ['yellow', 'red','green', 'oragne', 'purple'], 'Places': ['tokyo','bejing','washington','mumbai'], } # Words -> category categories = {word: key for key, words in data.items() for word in words} # Load the whole embedding matrix embeddings_index = {} with open('glove.6B.100d.txt', encoding='utf-8') as f: for line in f: values = line.split() word = values[0] embed = np.array(values[1:], dtype=np.float32) embeddings_index[word] = embed print('Loaded %s word vectors.' % len(embeddings_index)) # Embeddings for available words data_embeddings = {key: value for key, value in embeddings_index.items() if key in categories.keys()} # Processing the query def process(query): query_embed = embeddings_index[query] scores = {} for word, embed in data_embeddings.items(): category = categories[word] dist = query_embed.dot(embed) dist /= len(data[category]) scores[category] = scores.get(category, 0) + dist return scores # Testing print(process('jonny')) print(process('green')) print(process('park'))
运行输出
Loaded 400000 word vectors. {'Names': 7.965438079833984, 'Places': -0.3282392770051956, 'Colors': 1.803783965110779} {'Names': 11.360316085815429, 'Places': 3.536876901984215, 'Colors': 21.82199630737305} {'Names': 10.234728145599364, 'Places': 8.739515662193298, 'Colors': 10.761297225952148}
修改需求及解决方案
1. 类别顺序不一致的原因及解决
原代码中scores是普通Python字典,其键的顺序由类别首次被添加到字典的顺序决定:data_embeddings是从GloVe文件加载的embeddings_index过滤而来,单词顺序完全跟随GloVe文件,因此类别首次出现的顺序和data定义的顺序无关。
如果要强制返回顺序和data一致,修改process函数,最后按data的键顺序重构结果:
def process(query): query_embed = embeddings_index[query] scores = {} for word, embed in data_embeddings.items(): category = categories[word] dist = query_embed.dot(embed) dist /= len(data[category]) scores[category] = scores.get(category, 0) + dist # 按data定义的类别顺序返回结果 return {key: scores[key] for key in data.keys()}
2. 从文本文件读取查询列表
假设TEST.txt每行存储一个查询单词,替换原测试代码为以下内容:
# 读取查询列表 def load_queries(file_path): with open(file_path, 'r', encoding='utf-8') as f: # 过滤空行,返回非空的查询词 queries = [line.strip() for line in f if line.strip()] return queries # 批量处理所有查询 queries = load_queries('TEST.txt') results = {} for query in queries: # 处理单词不在词向量库的情况 if query not in embeddings_index: results[query] = "无可用词向量" continue results[query] = process(query)
3. 将结果保存为JSON或CSV文件
保存为JSON文件
使用Python内置json模块实现:
import json # 保存为同名JSON文件 output_json = 'TEST.json' with open(output_json, 'w', encoding='utf-8') as f: json.dump(results, f, ensure_ascii=False, indent=2)
保存为CSV文件
使用Python内置csv模块,将每个查询的类别分数展开为列:
import csv # 保存为同名CSV文件 output_csv = 'TEST.csv' category_names = list(data.keys()) with open(output_csv, 'w', newline='', encoding='utf-8') as f: writer = csv.writer(f) # 写入表头:查询词 + 各个类别 writer.writerow(['查询词'] + category_names) # 逐行写入查询结果 for query, res in results.items(): if isinstance(res, str): # 无词向量的情况,对应列填提示信息 row = [query] + ['无']*len(category_names) else: row = [query] + [res[cat] for cat in category_names] writer.writerow(row)
内容的提问来源于stack exchange,提问作者SunnyD
相关产品推荐
相关产品推荐

