()
| 54 | |
| 55 | |
| 56 | def main() -> None: |
| 57 | parser = argparse.ArgumentParser() |
| 58 | |
| 59 | # data configuration |
| 60 | parser.add_argument("--input_path", type=str, required=True) |
| 61 | parser.add_argument("--output_path", type=str, required=True) |
| 62 | parser.add_argument("--prompt", type=str, required=True) |
| 63 | |
| 64 | # SPHINX model configuration |
| 65 | parser.add_argument("--sphinx_type", type=str, choices=["SPHINX", "SPHINX-1k"]) |
| 66 | parser.add_argument("--tokenizer_path", type=str) |
| 67 | parser.add_argument("--pretrained_path", type=str) |
| 68 | parser.add_argument("--model_parallel_size", type=int, choices=[1,2]) |
| 69 | |
| 70 | # generation configuration |
| 71 | parser.add_argument("--max_gen_len", type=int, default=1024) |
| 72 | parser.add_argument("--temperature", type=float, default=0.1) |
| 73 | parser.add_argument("--top_p", type=float, default=0.75) |
| 74 | |
| 75 | args = parser.parse_args() |
| 76 | |
| 77 | if args.sphinx_type == "SPHINX-1k": |
| 78 | args.llama_type = "llama_ens5" # SPHINX-1k |
| 79 | elif args.sphinx_type == "SPHINX": |
| 80 | args.llama_type = "llama_ens" |
| 81 | |
| 82 | misc.init_distributed_mode(args) |
| 83 | fs_init.initialize_model_parallel(args.model_parallel_size) |
| 84 | |
| 85 | with default_tensor_type(dtype=torch.bfloat16, device="cuda"): |
| 86 | model = MetaModel( |
| 87 | args.llama_type, llama_config=[], tokenizer_path=args.tokenizer_path, |
| 88 | with_visual=True, max_seq_len=4096, |
| 89 | ) |
| 90 | print("Loading pretrained weights ...") |
| 91 | load_result = load_tensor_parallel_model_list(model, [args.pretrained_path]) |
| 92 | print("load result:\n", load_result) |
| 93 | assert load_result == {'missing_keys': [], 'unexpected_keys': []}, "checkpoint and model mismatch" |
| 94 | model.eval() |
| 95 | |
| 96 | dataset = Dataset(getattr(model.llma, 'image_size', 224), args.input_path) # 448 for SPHINX-1k, 224 for SPHINX |
| 97 | dataloader = torch.utils.data.DataLoader( |
| 98 | dataset, batch_size=10, shuffle=False, num_workers=4, pin_memory=True, |
| 99 | sampler=get_local_indices( |
| 100 | fs_init.get_data_parallel_rank(), |
| 101 | fs_init.get_data_parallel_world_size(), |
| 102 | len(dataset), |
| 103 | ), |
| 104 | ) |
| 105 | |
| 106 | conv = default_conversation() |
| 107 | conv.load_qas([[args.prompt, None]]) |
| 108 | prompt = conv.get_prompt() |
| 109 | conv_sep = conv.response_end_signal |
| 110 | |
| 111 | |
| 112 | if dist.get_rank() == 0: |
| 113 | print("Formatted prompt:", repr(prompt)) |
no test coverage detected