MCPcopy Create free account
hub / github.com/Alpha-VLLM/LLaMA2-Accessory / main

Function main

SPHINX/batch_inference.py:56–160  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

54
55
56def 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))

Callers 1

batch_inference.pyFile · 0.70

Calls 9

generateMethod · 0.95
default_tensor_typeClass · 0.90
MetaModelClass · 0.90
printFunction · 0.85
DatasetClass · 0.85
load_qasMethod · 0.80
get_promptMethod · 0.80
get_local_indicesFunction · 0.70

Tested by

no test coverage detected