(raw_text)
| 65 | raise ValueError(f'unknown strategy {args.sampling_strategy}') |
| 66 | |
| 67 | def process(raw_text): |
| 68 | if args.with_id: |
| 69 | query_id, raw_text = raw_text.split('\t') |
| 70 | # add MASK |
| 71 | generation_mask = '[gMASK]' if args.task_mask else '[MASK]' |
| 72 | if 'MASK]' not in raw_text: |
| 73 | raw_text += ' ' + generation_mask |
| 74 | seq = tokenizer.EncodeAsIds(raw_text).tokenization |
| 75 | seq = [tokenizer.get_command('ENC').Id] + seq |
| 76 | if not raw_text.endswith('MASK]'): |
| 77 | seq = seq + [tokenizer.get_command('eos').Id] |
| 78 | print('raw text: {}\n'.format(raw_text)) |
| 79 | if len(seq) > args.max_sequence_length: |
| 80 | raise ValueError('text too long.') |
| 81 | |
| 82 | # generation |
| 83 | mbz = args.max_inference_batch_size |
| 84 | assert args.batch_size < mbz or args.batch_size % mbz == 0 |
| 85 | output_list = [seq] |
| 86 | # continually detect the first mark position |
| 87 | while True: |
| 88 | seq = output_list[0] # TODO find the best one |
| 89 | # detect |
| 90 | mask_tokens = ['MASK', 'sMASK', 'gMASK'] if args.task_mask else ['MASK'] |
| 91 | mask_tokens = [tokenizer.get_command(token).Id for token in mask_tokens] |
| 92 | mask_position = len(seq) |
| 93 | for token in mask_tokens: |
| 94 | try: |
| 95 | mask_position = min(mask_position, seq.index(token)) |
| 96 | except ValueError: |
| 97 | pass |
| 98 | if mask_position == len(seq): |
| 99 | break |
| 100 | |
| 101 | get_func = partial(get_masks_and_position_ids_glm, mask_position=mask_position, context_length=len(seq)) |
| 102 | output_list = [] |
| 103 | for tim in range(max(args.batch_size // mbz, 1)): |
| 104 | input_seq = torch.cuda.LongTensor( |
| 105 | seq + [tokenizer.get_command('sop').Id] + [-1] * (args.out_seq_length - len(seq) - 1), |
| 106 | device=args.device) |
| 107 | output = filling_sequence(model, input_seq, |
| 108 | batch_size=min(args.batch_size, mbz), |
| 109 | strategy=strategy, |
| 110 | log_attention_weights=None, |
| 111 | get_masks_and_position_ids=get_func |
| 112 | )[0] # we don't use mems, fill back |
| 113 | if isinstance(output, torch.Tensor): # different strategies |
| 114 | output = list(output) |
| 115 | |
| 116 | output_list.extend(output) |
| 117 | |
| 118 | # clip -1s and fill back generated things into seq |
| 119 | for i in range(len(output_list)): |
| 120 | output = output_list[i].tolist() |
| 121 | try: |
| 122 | unfinished = output.index(-1) |
| 123 | except ValueError: |
| 124 | unfinished = len(output) |
nothing calls this directly
no test coverage detected