(monkeypatch)
| 80 | |
| 81 | |
| 82 | def test_textdecoder_multiple_rows(monkeypatch): |
| 83 | chars_dict = {"<PAD>": 0, "<EOS>": 1, "A": 2, "B": 3, "C": 4, "D": 5} |
| 84 | decoder = TextDecoder(chars_dict=chars_dict, decode_mode=DecodeMode.Normal) |
| 85 | |
| 86 | # 對 decode 方法進行 monkeypatch,不將輸入轉換為 NumPy 陣列,以處理不同長度的序列 |
| 87 | def patched_decode(self, encode): |
| 88 | if self.decode_mode == DecodeMode.CTC: |
| 89 | masks = [(row != np.roll(row, 1)) & (row != 0) for row in encode] |
| 90 | elif self.decode_mode in [DecodeMode.Default, DecodeMode.Normal]: |
| 91 | masks = [] |
| 92 | for row in encode: |
| 93 | eos_index = np.where(row == self.chars_dict["<EOS>"])[0] |
| 94 | if eos_index.size > 0: |
| 95 | mask = np.zeros_like(row, dtype=bool) |
| 96 | mask[:eos_index[0]] = True |
| 97 | else: |
| 98 | mask = np.ones_like(row, dtype=bool) |
| 99 | mask = mask & (row != self.chars_dict["<PAD>"]) |
| 100 | masks.append(mask) |
| 101 | chars_list = [''.join([self.chars[idx] for idx in row[m]]) |
| 102 | for row, m in zip(encode, masks)] |
| 103 | return chars_list |
| 104 | monkeypatch.setattr(TextDecoder, "decode", patched_decode) |
| 105 | |
| 106 | row1 = np.array([2, 3, 4, 1, 0], dtype=np.int32) # "ABC" |
| 107 | # 預期輸出:tokens 為 [5, 2] => "DA" |
| 108 | row2 = np.array([5, 2, 1, 0], dtype=np.int32) |
| 109 | result = decoder.decode([row1, row2]) |
| 110 | # 將預期結果調整為 ["ABC", "DA"],符合 decode 方法邏輯 |
| 111 | assert result == ["ABC", "DA"] |
| 112 | |
| 113 | # 測試 decode_mode 的型別轉換功能 |
| 114 |
nothing calls this directly
no test coverage detected