viterbi.py 1.7 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152
  1. # -*- coding: utf-8 -*-
  2. """Viterbi 解码算法。
  3. 按 ARCHITECTURE.md Phase 3 拆分建议,从 ``common/Utils.py`` 迁出。
  4. 类型:CORE(模型推理后处理,人工主导)。
  5. 原位置:``common/Utils.py`` 第 129-157 行的 ``viterbi_decode``。
  6. ``common/Utils.py`` 仍 re-export ``viterbi_decode``,老 import 不受影响。
  7. 说明:
  8. ARCHITECTURE.md 提到 ``decode``、``viterbi_decode``。``common/Utils.py``
  9. 中只有 ``viterbi_decode``;名为 ``decode`` 的函数分散在 ``foolnltk/``、
  10. ``product/``、``complaint/`` 等模块中,不属于 Utils.py 拆分范围。
  11. 本文件只迁 ``viterbi_decode``。
  12. """
  13. from __future__ import absolute_import
  14. import numpy as np
  15. __all__ = ["viterbi_decode"]
  16. def viterbi_decode(score, transition_params):
  17. """Decode the highest scoring sequence of tags outside of TensorFlow.
  18. This should only be used at test time.
  19. Args:
  20. score: A [seq_len, num_tags] matrix of unary potentials.
  21. transition_params: A [num_tags, num_tags] matrix of binary potentials.
  22. Returns:
  23. viterbi: A [seq_len] list of integers containing the highest scoring tag
  24. indices.
  25. viterbi_score: A float containing the score for the Viterbi sequence.
  26. """
  27. trellis = np.zeros_like(score)
  28. backpointers = np.zeros_like(score, dtype=np.int32)
  29. trellis[0] = score[0]
  30. for t in range(1, score.shape[0]):
  31. v = np.expand_dims(trellis[t - 1], 1) + transition_params
  32. trellis[t] = score[t] + np.max(v, 0)
  33. backpointers[t] = np.argmax(v, 0)
  34. viterbi = [np.argmax(trellis[-1])]
  35. for bp in reversed(backpointers[1:]):
  36. viterbi.append(bp[viterbi[-1]])
  37. viterbi.reverse()
  38. viterbi_score = np.max(trellis[-1])
  39. return viterbi, viterbi_score