Skip to content

better reporting of errors when using TP #14533

Description

@sayakpaul

I think the following aren't implemented at the moment (which is fine; just flagging).

  • Sharded loading. from_pretrained should stream shards straight to each rank's DTensor rather than materializing the full checkpoint then slicing — otherwise TP saves you nothing at load time. And check save_pretrained / state_dict calls .full_tensor() or uses DCP. I think we should at least raise when save_pretrained() is called in case TP is enabled?
  • LoRA loading. for a colwise base layer, lora_A replicated + lora_B colwise; for rowwise, lora_A rowwise + lora_B replicated. If the plan doesn't cover PEFT layers, loading an adapter onto a TP model will either error or be wrong. I think we should detect if the model has peft layers injected and raise if TP is requested?
  • Quantization, offloading. We should probably also raise when these are requested?

Originally posted by @sayakpaul in #13718 (comment)

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

No labels
No labels

Type

No type

Projects

Relationships

None yet

Development

No branches or pull requests

Issue actions