embed.py 4.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193
  1. # -*- coding: utf-8 -*-
  2. """词向量 / 字向量加载与 embedding 查找。
  3. 按 ARCHITECTURE.md Phase 3 拆分建议,从 ``common/Utils.py`` 迁出。
  4. 类型:CORE(模型推理基础,人工主导)。
  5. 原位置:``common/Utils.py`` 中以下函数和全局变量:
  6. - ``model_w2v`` / ``lock_model_w2v`` / ``model_word`` / ``lock_model_word``
  7. - ``model_word_file``
  8. - ``getw2vfilepath`` / ``getFileFromSysPath``
  9. - ``getModel_w2v`` / ``getModel_word``
  10. - ``embedding`` / ``embedding_word`` / ``embedding_word_forward``
  11. - ``formEncoding``
  12. 本文件自包含,不依赖 ``common/Utils.py``,避免循环 import。
  13. ``common/Utils.py`` 仍 re-export 以上全部名称,老 import 不受影响。
  14. 按 ARCHITECTURE.md §4.3 依赖方向约束:
  15. model_runtime -> domain, infra only
  16. """
  17. from __future__ import absolute_import
  18. import os
  19. import sys
  20. import re
  21. from threading import RLock
  22. import numpy as np
  23. import gensim
  24. __all__ = [
  25. "model_w2v",
  26. "lock_model_w2v",
  27. "model_word",
  28. "lock_model_word",
  29. "model_word_file",
  30. "getw2vfilepath",
  31. "getFileFromSysPath",
  32. "getModel_w2v",
  33. "getModel_word",
  34. "embedding",
  35. "embedding_word",
  36. "embedding_word_forward",
  37. "formEncoding",
  38. ]
  39. model_w2v = None
  40. lock_model_w2v = RLock()
  41. model_word_file = os.path.dirname(os.path.abspath(__file__)) + "/../singlew2v_model.vector"
  42. model_word = None
  43. lock_model_word = RLock()
  44. Lazy_load = False
  45. def getLazyLoad():
  46. global Lazy_load
  47. return Lazy_load
  48. def getFileFromSysPath(filename):
  49. for _path in sys.path:
  50. if os.path.isdir(_path):
  51. for _file in os.listdir(_path):
  52. _abspath = os.path.join(_path, _file)
  53. if os.path.isfile(_abspath):
  54. if _file == filename:
  55. return _abspath
  56. return None
  57. def getw2vfilepath():
  58. filename = "wiki_128_word_embedding_new.vector"
  59. w2vfile = getFileFromSysPath(filename)
  60. if w2vfile is not None:
  61. return w2vfile
  62. return filename
  63. def getModel_w2v():
  64. '''
  65. @summary:加载词向量
  66. '''
  67. global model_w2v, lock_model_w2v
  68. with lock_model_w2v:
  69. if model_w2v is None:
  70. model_w2v = gensim.models.KeyedVectors.load_word2vec_format(getw2vfilepath(), binary=True)
  71. return model_w2v
  72. def getModel_word():
  73. '''
  74. @summary:加载字向量
  75. '''
  76. global model_word, lock_model_w2v
  77. with lock_model_word:
  78. if model_word is None:
  79. model_word = gensim.models.KeyedVectors.load_word2vec_format(model_word_file, binary=True)
  80. return model_word
  81. def embedding(datas, shape):
  82. '''
  83. @summary:查找词汇对应的词向量
  84. @param:
  85. datas:词汇的list
  86. shape:结果的shape
  87. @return: array,返回对应shape的词嵌入
  88. '''
  89. model_w2v = getModel_w2v()
  90. embed = np.zeros(shape)
  91. length = shape[1]
  92. out_index = 0
  93. for data in datas:
  94. index = 0
  95. for item in data:
  96. item_not_space = re.sub("\s*", "", item)
  97. if index >= length:
  98. break
  99. if item_not_space in model_w2v.vocab:
  100. embed[out_index][index] = model_w2v[item_not_space]
  101. index += 1
  102. else:
  103. index += 1
  104. out_index += 1
  105. return embed
  106. def embedding_word(datas, shape):
  107. '''
  108. @summary:查找词汇对应的词向量
  109. @param:
  110. datas:词汇的list
  111. shape:结果的shape
  112. @return: array,返回对应shape的词嵌入
  113. '''
  114. model_w2v = getModel_word()
  115. embed = np.zeros(shape)
  116. length = shape[1]
  117. out_index = 0
  118. for data in datas:
  119. index = 0
  120. for item in str(data)[-shape[1]:]:
  121. if index >= length:
  122. break
  123. if item in model_w2v.vocab:
  124. embed[out_index][index] = model_w2v[item]
  125. index += 1
  126. else:
  127. index += 1
  128. out_index += 1
  129. return embed
  130. def embedding_word_forward(datas, shape):
  131. '''
  132. @summary:查找词汇对应的词向量
  133. @param:
  134. datas:词汇的list
  135. shape:结果的shape
  136. @return: array,返回对应shape的词嵌入
  137. '''
  138. model_w2v = getModel_word()
  139. embed = np.zeros(shape)
  140. length = shape[1]
  141. out_index = 0
  142. for data in datas:
  143. index = 0
  144. for item in str(data)[:shape[1]]:
  145. if index >= length:
  146. break
  147. if item in model_w2v.vocab:
  148. embed[out_index][index] = model_w2v[item]
  149. index += 1
  150. else:
  151. index += 1
  152. out_index += 1
  153. return embed
  154. def formEncoding(text, shape=(100, 60), expand=False):
  155. embedding = np.zeros(shape)
  156. word_model = getModel_word()
  157. for i in range(len(text)):
  158. if i >= shape[0]:
  159. break
  160. if text[i] in word_model.vocab:
  161. embedding[i] = word_model[text[i]]
  162. if expand:
  163. embedding = np.expand_dims(embedding, 0)
  164. return embedding