I would like to know how fast this is compared to a model shared on 3 A100s with 32 GBs of VRAM each.
What I am also interested in, is why model sharding has to be done manually. It seems like, one should be able to write a framework that will take your forward step and distribute the amount of layers on the available GPUs, automatically. But I haven't come across such a framework yet.