(rank, world_size, data, max_new_tokens, fout, template, cache_fout, cache_dict)
| 54 | return '' |
| 55 | |
| 56 | def get_pred(rank, world_size, data, max_new_tokens, fout, template, cache_fout, cache_dict): |
| 57 | for item in tqdm(data): |
| 58 | try: |
| 59 | inst = item['prompt'] |
| 60 | plan = item['plan'].strip().replace('\n\n', '\n') |
| 61 | steps = plan.split('\n') |
| 62 | text = "" |
| 63 | responses = [] |
| 64 | if len(steps) > 50: |
| 65 | print(plan) |
| 66 | continue |
| 67 | for step in steps: |
| 68 | if inst in cache_dict and step in cache_dict[inst]: |
| 69 | response = cache_dict[inst][step] |
| 70 | responses.append(response) |
| 71 | text += response + '\n\n' |
| 72 | continue |
| 73 | prompt = template.replace('$INST$', inst).replace('$PLAN$', plan.strip()).replace('$TEXT$', text.strip()).replace('$STEP$', step.strip()) |
| 74 | response = get_response_gpt4(prompt, max_new_tokens) |
| 75 | if response == '': |
| 76 | break |
| 77 | # save to cache |
| 78 | cache_fout.write(json.dumps({"prompt": inst, "step": step, "response": response}, ensure_ascii=False)+'\n') |
| 79 | cache_fout.flush() |
| 80 | responses.append(response) |
| 81 | text += response + '\n\n' |
| 82 | if response == '': |
| 83 | continue |
| 84 | item["write"] = responses |
| 85 | fout.write(json.dumps(item, ensure_ascii=False)+'\n') |
| 86 | fout.flush() |
| 87 | except Exception as e: |
| 88 | print(e) |
| 89 | |
| 90 | def seed_everything(seed): |
| 91 | torch.manual_seed(seed) |
nothing calls this directly
no test coverage detected