Trainer के साथ ग्रेडिएंट चेकपॉइंटिंग
आप अपने मॉडल की मेमोरी खपत कम करने के लिए ग्रेडिएंट चेकपॉइंटिंग का उपयोग करना चाहते हैं. आपने Accelerator के साथ explicit training loop लिखना देखा है, और अब आप Trainer के साथ बिना training loops वाला सरल इंटरफेस उपयोग करना चाहते हैं. trainer.train() कॉल के कारण यह अभ्यास चलने में कुछ समय लेगा.
ग्रेडिएंट चेकपॉइंटिंग इस्तेमाल करने के लिए Trainer के आर्ग्युमेंट सेट करें.
यह अभ्यास पाठ्यक्रम का हिस्सा है
PyTorch के साथ कुशल AI मॉडल प्रशिक्षण
अभ्यास निर्देश
TrainingArgumentsमें चार gradient accumulation steps उपयोग करें.TrainingArgumentsमें gradient checkpointing सक्षम करें.- Training arguments को
Trainerमें पास करें.
इंटरैक्टिव व्यावहारिक अभ्यास
इस अभ्यास को इस नमूना कोड को पूरा करके आज़माएँ।
training_args = TrainingArguments(output_dir="./results",
evaluation_strategy="epoch",
# Use four gradient accumulation steps
gradient_accumulation_steps=____,
# Enable gradient checkpointing
gradient_checkpointing=____)
trainer = Trainer(model=model,
# Pass in the training arguments
args=____,
train_dataset=dataset["train"],
eval_dataset=dataset["validation"],
compute_metrics=compute_metrics)
trainer.train()