MCPcopy Create free account
hub / github.com/TIGER-AI-Lab/ScholarCopilot / main

Function main

train/src/encode.py:28–120  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

26
27
28def main():
29 parser = HfArgumentParser((ModelArguments, DataArguments, TrainingArguments))
30 if len(sys.argv) == 2 and sys.argv[1].endswith(".json"):
31 model_args, data_args, training_args = parser.parse_json_file(json_file=os.path.abspath(sys.argv[1]))
32 else:
33 model_args, data_args, training_args = parser.parse_args_into_dataclasses()
34 model_args: ModelArguments
35 data_args: DataArguments
36 training_args: TrainingArguments
37
38 if training_args.local_rank > 0 or training_args.n_gpu > 1:
39 raise NotImplementedError('Multi-GPU encoding is not supported.')
40
41 # Setup logging
42 logging.basicConfig(
43 format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
44 datefmt="%m/%d/%Y %H:%M:%S",
45 level=logging.INFO if training_args.local_rank in [-1, 0] else logging.WARN,
46 )
47 tokenizer = AutoTokenizer.from_pretrained(
48 model_args.tokenizer_name if model_args.tokenizer_name else model_args.model_name_or_path,
49 cache_dir=model_args.cache_dir
50 )
51 if tokenizer.pad_token_id is None:
52 tokenizer.pad_token_id = tokenizer.eos_token_id
53 tokenizer.padding_side = 'right'
54 tokenizer.add_tokens(['<|paper_start|>', '<|paper_end|>', '<|cite_start|>', '<|cite_end|>', '<|reference_start|>', '<|reference_end|>'])
55
56 if training_args.bf16:
57 torch_dtype = torch.bfloat16
58 elif training_args.fp16:
59 torch_dtype = torch.float16
60 else:
61 torch_dtype = torch.float32
62
63 model = ArxivLLM.load(
64 model_args.model_name_or_path,
65 pooling=model_args.pooling,
66 normalize=model_args.normalize,
67 lora_name_or_path=model_args.lora_name_or_path,
68 cache_dir=model_args.cache_dir,
69 torch_dtype=torch_dtype
70 )
71
72 model.encoder.resize_token_embeddings(len(tokenizer))
73
74 encode_dataset = EncodeDataset(
75 data_args=data_args,
76 )
77
78 encode_collator = EncodeCollator(
79 data_args=data_args,
80 tokenizer=tokenizer,
81 )
82
83 encode_loader = DataLoader(
84 encode_dataset,
85 batch_size=training_args.per_device_eval_batch_size,

Callers 1

encode.pyFile · 0.70

Calls 3

EncodeDatasetClass · 0.90
EncodeCollatorClass · 0.90
loadMethod · 0.80

Tested by

no test coverage detected