MCPcopy Create free account
hub / github.com/AnswerDotAI/ModernBERT / train

Function train

sequence_classification.py:243–331  ·  view source on GitHub ↗
(cfg: DictConfig, return_trainer: bool = False, do_train: bool = True)

Source from the content-addressed store, hash-verified

241
242
243def train(cfg: DictConfig, return_trainer: bool = False, do_train: bool = True) -> Optional[Trainer]:
244 print("Training using config: ")
245 print(om.to_yaml(cfg))
246 reproducibility.seed_all(cfg.seed)
247
248 # Get batch size info
249 cfg = update_batch_size_info(cfg)
250
251 # Build Model
252 print("Initializing model...")
253 model = build_model(cfg.model)
254 n_params = sum(p.numel() for p in model.parameters())
255 print(f"{n_params=:.4e}")
256
257 # Dataloaders
258 print("Building train loader...")
259 train_loader = build_my_dataloader(
260 cfg.train_loader,
261 cfg.global_train_batch_size // dist.get_world_size(),
262 )
263 print("Building eval loader...")
264 global_eval_batch_size = cfg.get("global_eval_batch_size", cfg.global_train_batch_size)
265 eval_loader = build_my_dataloader(
266 cfg.eval_loader,
267 cfg.get("device_eval_batch_size", global_eval_batch_size // dist.get_world_size()),
268 )
269 eval_evaluator = Evaluator(
270 label="eval",
271 dataloader=eval_loader,
272 device_eval_microbatch_size=cfg.get("device_eval_microbatch_size", None),
273 )
274
275 # Optimizer
276 optimizer = build_optimizer(cfg.optimizer, model)
277
278 # Scheduler
279 scheduler = build_scheduler(cfg.scheduler)
280
281 # Loggers
282 loggers = [build_logger(name, logger_cfg) for name, logger_cfg in cfg.get("loggers", {}).items()]
283
284 # Callbacks
285 callbacks = [build_callback(name, callback_cfg) for name, callback_cfg in cfg.get("callbacks", {}).items()]
286
287 # Algorithms
288 algorithms = [build_algorithm(name, algorithm_cfg) for name, algorithm_cfg in cfg.get("algorithms", {}).items()]
289
290 if cfg.get("run_name") is None:
291 cfg.run_name = os.environ.get("COMPOSER_RUN_NAME", "sequence-classification")
292
293 # Build the Trainer
294 trainer = Trainer(
295 run_name=cfg.run_name,
296 seed=cfg.seed,
297 model=model,
298 algorithms=algorithms,
299 train_dataloader=train_loader,
300 eval_dataloader=eval_evaluator,

Callers 1

Calls 9

build_my_dataloaderFunction · 0.85
update_batch_size_infoFunction · 0.70
build_modelFunction · 0.70
build_optimizerFunction · 0.70
build_schedulerFunction · 0.70
build_loggerFunction · 0.70
build_callbackFunction · 0.70
build_algorithmFunction · 0.70
log_configFunction · 0.70

Tested by

no test coverage detected