vocab.py 4.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129
  1. # -*- coding: utf-8 -*-
  2. """字/词 vocab 与 char/id 映射。
  3. 按 ARCHITECTURE.md Phase 3 拆分建议,从 ``common/Utils.py`` 迁出。
  4. 类型:CORE(模型推理基础,人工主导)。
  5. 原位置:``common/Utils.py`` 中以下函数和全局变量:
  6. - ``vocab_word`` / ``vocab_words``
  7. - ``file_vocab_word`` / ``file_vocab_words``
  8. - ``fool_char_to_id``(模块加载时从 ``fool_char_to_id.pk`` 读取)
  9. - ``getIndexOfWord`` / ``getIndexOfWords`` / ``getIndexOfWord_fool``
  10. - ``getVocabAndMatrix``
  11. - ``changeIndexFromWordToWords``
  12. 依赖关系:
  13. - ``getIndexOfWord`` / ``getIndexOfWords`` 调用 ``getModel_word`` / ``getModel_w2v``
  14. (来自 ``model_runtime.embed``)和 ``save`` / ``load``
  15. (仍留在 ``common/Utils.py``,本文件用延迟 import 避免循环依赖)。
  16. - ``fool_char_to_id`` 在模块加载时用 ``pickle.load`` 直接读取,不依赖 ``save`` / ``load``。
  17. ``common/Utils.py`` 仍 re-export 以上全部名称,老 import 不受影响。
  18. 按 ARCHITECTURE.md §4.3 依赖方向约束:
  19. model_runtime -> domain, infra only
  20. """
  21. from __future__ import absolute_import
  22. import os
  23. import pickle
  24. import numpy as np
  25. from BiddingKG.dl.model_runtime import embed as _embed
  26. __all__ = [
  27. "vocab_word",
  28. "vocab_words",
  29. "file_vocab_word",
  30. "file_vocab_words",
  31. "fool_char_to_id",
  32. "getIndexOfWord",
  33. "getIndexOfWords",
  34. "getIndexOfWord_fool",
  35. "getVocabAndMatrix",
  36. "changeIndexFromWordToWords",
  37. ]
  38. vocab_word = None
  39. vocab_words = None
  40. file_vocab_word = "vocab_word.pk"
  41. file_vocab_words = "vocab_words.pk"
  42. fool_char_to_id = pickle.load(
  43. open(os.path.dirname(os.path.abspath(__file__)) + "/../common/fool_char_to_id.pk", 'rb')
  44. )
  45. def getVocabAndMatrix(model, Embedding_size=60):
  46. '''
  47. @summary:获取子向量的词典和子向量矩阵
  48. '''
  49. vocab = ["<pad>"] + model.index2word
  50. embedding_matrix = np.zeros((len(vocab), Embedding_size))
  51. for i in range(1, len(vocab)):
  52. embedding_matrix[i] = model[vocab[i]]
  53. return vocab, embedding_matrix
  54. def getIndexOfWord(word):
  55. global vocab_word, file_vocab_word
  56. if vocab_word is None:
  57. if os.path.exists(file_vocab_word):
  58. # 延迟 import save/load 避免 common/Utils.py <-> model_runtime/vocab.py 循环依赖
  59. from BiddingKG.dl.common.Utils import load, save
  60. vocab = load(file_vocab_word)
  61. vocab_word = dict((w, i) for i, w in enumerate(np.array(vocab)))
  62. else:
  63. model = _embed.getModel_word()
  64. vocab, _ = getVocabAndMatrix(model, Embedding_size=60)
  65. vocab_word = dict((w, i) for i, w in enumerate(np.array(vocab)))
  66. from BiddingKG.dl.common.Utils import save
  67. save(vocab, file_vocab_word)
  68. if word in vocab_word.keys():
  69. return vocab_word[word]
  70. else:
  71. return vocab_word['<pad>']
  72. def changeIndexFromWordToWords(tokens, word_index):
  73. '''
  74. @summary:转换某个字的字偏移为词偏移
  75. '''
  76. before_index = 0
  77. after_index = 0
  78. for i in range(len(tokens)):
  79. after_index = after_index + len(tokens[i])
  80. if before_index <= word_index and after_index > word_index:
  81. return i
  82. before_index = after_index
  83. return i + 1
  84. def getIndexOfWords(words):
  85. global vocab_words, file_vocab_words
  86. if vocab_words is None:
  87. if os.path.exists(file_vocab_words):
  88. from BiddingKG.dl.common.Utils import load, save
  89. vocab = load(file_vocab_words)
  90. vocab_words = dict((w, i) for i, w in enumerate(np.array(vocab)))
  91. else:
  92. model = _embed.getModel_w2v()
  93. vocab, _ = getVocabAndMatrix(model, Embedding_size=128)
  94. vocab_words = dict((w, i) for i, w in enumerate(np.array(vocab)))
  95. from BiddingKG.dl.common.Utils import save
  96. save(vocab, file_vocab_words)
  97. if words in vocab_words.keys():
  98. return vocab_words[words]
  99. else:
  100. return vocab_words["<pad>"]
  101. def getIndexOfWord_fool(word):
  102. if word in fool_char_to_id.keys():
  103. return fool_char_to_id[word]
  104. else:
  105. return fool_char_to_id["[UNK]"]