MCPcopy Create free account
hub / github.com/OpenBMB/AgentCPM-GUI / make_supervised_data_module

Function make_supervised_data_module

sft/finetune.py:86–144  ·  view source on GitHub ↗

Make dataset and collator for supervised fine-tuning.

(
    tokenizer: transformers.PreTrainedTokenizer,
    data_args,
    transform,
    data_collator=None,
    llm_type="minicpm",
    slice_config=None,
    patch_size=14,
    query_nums=64,
    batch_vision=False,
    max_length=2048,
)

Source from the content-addressed store, hash-verified

84
85
86def make_supervised_data_module(
87 tokenizer: transformers.PreTrainedTokenizer,
88 data_args,
89 transform,
90 data_collator=None,
91 llm_type="minicpm",
92 slice_config=None,
93 patch_size=14,
94 query_nums=64,
95 batch_vision=False,
96 max_length=2048,
97) -> Dict:
98 """Make dataset and collator for supervised fine-tuning."""
99 dataset_cls = SupervisedDataset
100
101 rank0_print("Loading data...")
102
103 def load(path):
104 if not path.endswith(('.json', '.jsonl')):
105 raise ValueError('need .json or .jsonl')
106 with open(path, encoding='utf-8') as f:
107 return json.load(f) if path.endswith('.json') else [json.loads(l) for l in f if l.strip()]
108 train_json = load(data_args.data_path)
109
110 train_dataset = dataset_cls(
111 train_json,
112 transform,
113 tokenizer,
114 slice_config=slice_config,
115 llm_type=llm_type,
116 patch_size=patch_size,
117 query_nums=query_nums,
118 batch_vision=batch_vision,
119 max_length=max_length,
120 max_line_res=data_args.max_line_res
121 )
122
123 if data_args.eval_data_path:
124 eval_json = load(data_args.eval_data_path)
125 eval_dataset = dataset_cls(
126 eval_json,
127 transform,
128 tokenizer,
129 slice_config=slice_config,
130 llm_type=llm_type,
131 patch_size=patch_size,
132 query_nums=query_nums,
133 batch_vision=batch_vision,
134 max_length=max_length,
135 max_line_res=data_args.max_line_res
136 )
137 else:
138 eval_dataset = None
139
140 return dict(
141 train_dataset=train_dataset,
142 eval_dataset=eval_dataset,
143 data_collator= partial(data_collator, max_length=max_length),

Callers 1

trainFunction · 0.85

Calls 2

rank0_printFunction · 0.85
loadFunction · 0.85

Tested by

no test coverage detected