MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / build_train_valid_test_data_iterators

Function build_train_valid_test_data_iterators

codegeex/megatron/training.py:1220–1360  ·  view source on GitHub ↗
(build_train_valid_test_datasets_provider)

Source from the content-addressed store, hash-verified

1218
1219
1220def build_train_valid_test_data_iterators(build_train_valid_test_datasets_provider):
1221 args = get_args()
1222
1223 (train_dataloader, valid_dataloader, test_dataloader) = (None, None, None)
1224
1225 print_rank_0("> building train, validation, and test datasets ...")
1226
1227 # Backward compatibility, assume fixed batch size.
1228 if args.iteration > 0 and args.consumed_train_samples == 0:
1229 assert (
1230 args.train_samples is None
1231 ), "only backward compatibility support for iteration-based training"
1232 args.consumed_train_samples = args.iteration * args.global_batch_size
1233 if args.iteration > 0 and args.consumed_valid_samples == 0:
1234 assert (
1235 args.train_samples is None
1236 ), "only backward compatibility support for iteration-based training"
1237 args.consumed_valid_samples = (
1238 (args.iteration // args.eval_interval)
1239 * args.eval_iters
1240 * args.global_batch_size
1241 )
1242
1243 # Data loader only on rank 0 of each model parallel group.
1244 if mpu.get_tensor_model_parallel_rank() == 0:
1245
1246 # Number of train/valid/test samples.
1247 if args.train_samples:
1248 train_samples = args.train_samples
1249 else:
1250 train_samples = args.train_iters * args.global_batch_size
1251 eval_iters = (args.train_iters // args.eval_interval + 1) * args.eval_iters
1252 test_iters = args.eval_iters
1253 train_val_test_num_samples = [
1254 train_samples,
1255 eval_iters * args.global_batch_size,
1256 test_iters * args.global_batch_size,
1257 ]
1258 print_rank_0(" > datasets target sizes (minimum size):")
1259 print_rank_0(" train: {}".format(train_val_test_num_samples[0]))
1260 print_rank_0(" validation: {}".format(train_val_test_num_samples[1]))
1261 print_rank_0(" test: {}".format(train_val_test_num_samples[2]))
1262
1263 # Build the datasets.
1264 train_ds, valid_ds, test_ds = build_train_valid_test_datasets_provider(
1265 train_val_test_num_samples
1266 )
1267
1268 # Build dataloders.
1269 train_dataloader = build_pretraining_data_loader(
1270 train_ds, args.consumed_train_samples
1271 )
1272 if args.co_evaluation:
1273 valid_dataloader = {}
1274 for key, value in valid_ds.items():
1275 valid_dataloader[key] = build_pretraining_data_loader(
1276 value, args.consumed_valid_samples
1277 )

Callers 1

pretrainFunction · 0.85

Calls 4

get_argsFunction · 0.90
print_rank_0Function · 0.90
cyclic_iterFunction · 0.85

Tested by

no test coverage detected