MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / process

Function process

SwissArmyTransformer/examples/glm/inference_glm.py:67–146  ·  view source on GitHub ↗
(raw_text)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 8

filling_sequenceFunction · 0.90
timed_nameFunction · 0.90
extendMethod · 0.80
appendMethod · 0.80
printFunction · 0.50
EncodeAsIdsMethod · 0.45
get_commandMethod · 0.45
DecodeIdsMethod · 0.45

Tested by

no test coverage detected