Gradient checkpointing และ Local SGD

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Dennis Lee

Data Engineer, Amazon

การเพิ่มประสิทธิภาพการฝึกอบรม

 

 

ไอคอนแสดงประสิทธิภาพด้านหน่วยความจำ การสื่อสาร และการคำนวณ

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing ช่วยเพิ่มประสิทธิภาพด้านหน่วยความจำ

 

 

ไอคอนแสดงประสิทธิภาพด้านหน่วยความจำ การสื่อสาร และการคำนวณ

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Local SGD แก้ปัญหาประสิทธิภาพด้านการสื่อสาร

 

 

ไอคอนแสดงประสิทธิภาพด้านหน่วยความจำ การสื่อสาร และการคำนวณ

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing คืออะไร?

  • Gradient checkpointing: ลดหน่วยความจำโดยเลือก activation ที่จะบันทึก
  • ตัวอย่าง: คำนวณ A + B = C

กราฟแสดง gradient checkpointing ด้วยโหนดและเส้นเชื่อม

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing คืออะไร?

  • Gradient checkpointing: ลดหน่วยความจำโดยเลือก activation ที่จะบันทึก
  • ตัวอย่าง: คำนวณ A + B = C
    • คำนวณ A, B ก่อน แล้วจึงคำนวณ C

กราฟแสดง gradient checkpointing ด้วยโหนดและเส้นเชื่อม

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing คืออะไร?

  • Gradient checkpointing: ลดหน่วยความจำโดยเลือก activation ที่จะบันทึก
  • ตัวอย่าง: คำนวณ A + B = C
    • คำนวณ A, B ก่อน แล้วจึงคำนวณ C
    • A, B ไม่จำเป็นในส่วนที่เหลือของ forward pass
  • ควรบันทึกหรือลบ A และ B?

กราฟแสดง gradient checkpointing ด้วยโหนดและเส้นเชื่อม

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing คืออะไร?

  • Gradient checkpointing: ลดหน่วยความจำโดยเลือก activation ที่จะบันทึก
  • ตัวอย่าง: คำนวณ A + B = C
    • คำนวณ A, B ก่อน แล้วจึงคำนวณ C
    • A, B ไม่จำเป็นในส่วนที่เหลือของ forward pass
  • ควรบันทึกหรือลบ A และ B?
    • ไม่ใช้ gradient checkpointing: บันทึก A, B

กราฟแสดง gradient checkpointing ด้วยโหนดและเส้นเชื่อม

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing คืออะไร?

  • Gradient checkpointing: ลดหน่วยความจำโดยเลือก activation ที่จะบันทึก
  • ตัวอย่าง: คำนวณ A + B = C
    • คำนวณ A, B ก่อน แล้วจึงคำนวณ C
    • A, B ไม่จำเป็นในส่วนที่เหลือของ forward pass
  • ควรบันทึกหรือลบ A และ B?
    • ไม่ใช้ gradient checkpointing: บันทึก A, B
    • ใช้ gradient checkpointing: ลบ A, B

กราฟแสดง gradient checkpointing ด้วยโหนดและเส้นเชื่อม

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing คืออะไร?

  • Gradient checkpointing: ลดหน่วยความจำโดยเลือก activation ที่จะบันทึก
  • ตัวอย่าง: คำนวณ A + B = C
    • คำนวณ A, B ก่อน แล้วจึงคำนวณ C
    • A, B ไม่จำเป็นในส่วนที่เหลือของ forward pass
  • ควรบันทึกหรือลบ A และ B?
    • ไม่ใช้ gradient checkpointing: บันทึก A, B
    • ใช้ gradient checkpointing: ลบ A, B
    • คำนวณ A, B ใหม่ระหว่าง backward pass

กราฟแสดง gradient checkpointing ด้วยโหนดและเส้นเชื่อม

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing คืออะไร?

  • Gradient checkpointing: ลดหน่วยความจำโดยเลือก activation ที่จะบันทึก
  • ตัวอย่าง: คำนวณ A + B = C
    • คำนวณ A, B ก่อน แล้วจึงคำนวณ C
    • A, B ไม่จำเป็นในส่วนที่เหลือของ forward pass
  • ควรบันทึกหรือลบ A และ B?
    • ไม่ใช้ gradient checkpointing: บันทึก A, B
    • ใช้ gradient checkpointing: ลบ A, B
    • คำนวณ A, B ใหม่ระหว่าง backward pass
    • หาก B ใช้ต้นทุนการคำนวณสูง ให้บันทึกไว้

กราฟแสดง gradient checkpointing ด้วยโหนดและเส้นเชื่อม

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Trainer และ Accelerator

แผนภูมิเปรียบเทียบความง่ายในการใช้งานกับความสามารถในการปรับแต่งของ Accelerator และ Trainer

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Trainer และ Accelerator

แผนภูมิเปรียบเทียบความง่ายในการใช้งานกับความสามารถในการปรับแต่งของ Accelerator และ Trainer

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing ด้วย Trainer

training_args = TrainingArguments(output_dir="./results",
                                  evaluation_strategy="epoch",
                                  gradient_accumulation_steps=4)







การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing ด้วย Trainer

training_args = TrainingArguments(output_dir="./results",
                                  evaluation_strategy="epoch",
                                  gradient_accumulation_steps=4,
                                  gradient_checkpointing=True)

trainer = Trainer(model=model, args=training_args, train_dataset=dataset["train"], eval_dataset=dataset["validation"], compute_metrics=compute_metrics)
trainer.train()
{'epoch': 1.0, 'eval_loss': 0.73, 'eval_accuracy': 0.03, 'eval_f1': 0.05}
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

จาก Trainer สู่ Accelerator

แผนภูมิเปรียบเทียบความง่ายในการใช้งานกับความสามารถในการปรับแต่งของ Accelerator และ Trainer

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing ด้วย Accelerator

accelerator = Accelerator(gradient_accumulation_steps=2)


for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        inputs, targets = batch["input_ids"], batch["labels"]
        outputs = model(inputs, labels=targets)
        loss = outputs.loss
        accelerator.backward(loss)
        optimizer.step()
        lr_scheduler.step()
        optimizer.zero_grad()
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Gradient checkpointing ด้วย Accelerator

accelerator = Accelerator(gradient_accumulation_steps=2)
model.gradient_checkpointing_enable()

for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        inputs, targets = batch["input_ids"], batch["labels"]
        outputs = model(inputs, labels=targets)
        loss = outputs.loss
        accelerator.backward(loss)
        optimizer.step()
        lr_scheduler.step()
        optimizer.zero_grad()
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Local SGD ช่วยเพิ่มประสิทธิภาพด้านการสื่อสาร

 

 

ไอคอนแสดงประสิทธิภาพด้านหน่วยความจำ การสื่อสาร และการคำนวณ

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Local SGD คืออะไร?

ไดอะแกรมแสดงการทำงานของ Local SGD โดยซิงโครไนซ์ gradient หลังจากจำนวนขั้นตอนที่กำหนด

  • แต่ละอุปกรณ์คำนวณ gradient แบบขนาน
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Local SGD คืออะไร?

ไดอะแกรมแสดงการทำงานของ Local SGD โดยซิงโครไนซ์ gradient หลังจากจำนวนขั้นตอนที่กำหนด

  • แต่ละอุปกรณ์คำนวณ gradient แบบขนาน
  • การซิงโครไนซ์ gradient: Driver node อัปเดตพารามิเตอร์โมเดลในแต่ละอุปกรณ์
  • Local SGD: ลดความถี่ในการซิงโครไนซ์ gradient
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Local SGD ด้วย Accelerator





for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        inputs, targets = batch["input_ids"], batch["labels"]
        outputs = model(inputs, labels=targets)
        loss = outputs.loss
        accelerator.backward(loss)
        optimizer.step()
        lr_scheduler.step()
        optimizer.zero_grad()

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Local SGD ด้วย Accelerator

from accelerate.local_sgd import LocalSGD

with LocalSGD(accelerator=accelerator, model=model, local_sgd_steps=8, 
              enabled=True) as local_sgd:
    for index, batch in enumerate(dataloader):
        with accelerator.accumulate(model):
            inputs, targets = batch["input_ids"], batch["labels"]
            outputs = model(inputs, labels=targets)
            loss = outputs.loss
            accelerator.backward(loss)
            optimizer.step()
            lr_scheduler.step()
            optimizer.zero_grad()

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Local SGD ด้วย Accelerator

from accelerate.local_sgd import LocalSGD

with LocalSGD(accelerator=accelerator, model=model, local_sgd_steps=8, 
              enabled=True) as local_sgd:
    for index, batch in enumerate(dataloader):
        with accelerator.accumulate(model):
            inputs, targets = batch["input_ids"], batch["labels"]
            outputs = model(inputs, labels=targets)
            loss = outputs.loss
            accelerator.backward(loss)
            optimizer.step()
            lr_scheduler.step()
            optimizer.zero_grad()
            local_sgd.step()
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

ฝึกปฏิบัติกันเลย!

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Preparing Video For Download...