| 115 | return [c + numres for c in s.encode("utf-8")] |
| 116 | |
| 117 | def decode(self, ids, strip_extraneous=False): |
| 118 | if strip_extraneous: |
| 119 | ids = strip_ids(ids, list(range(self._num_reserved_ids or 0))) |
| 120 | numres = self._num_reserved_ids |
| 121 | decoded_ids = [] |
| 122 | int2byte = six.int2byte |
| 123 | for id_ in ids: |
| 124 | if 0 <= id_ < numres: |
| 125 | decoded_ids.append(RESERVED_TOKENS_BYTES[int(id_)]) |
| 126 | else: |
| 127 | decoded_ids.append(int2byte(id_ - numres)) |
| 128 | if six.PY2: |
| 129 | return "".join(decoded_ids) |
| 130 | # Python3: join byte arrays and then decode string |
| 131 | return b"".join(decoded_ids).decode("utf-8", "replace") |
| 132 | |
| 133 | def decode_list(self, ids): |
| 134 | numres = self._num_reserved_ids |