已开启
model parallize metadata #327
townwish4gitcode创建于 8月11日
townwish4gitcode
8月11日 评论:
8月11日 评论:
TODO: tp_grad_info实质上可以被包含于model_conversion_metadata,在此情形下ParamLayout应当扩展哪些字段


8月11日 修改了issue 的描述
8月11日 修改了issue 的描述
TODO: tp_grad_info实质上可以被包含于model_conversion_metadata,在此情形下ParamLayout应当扩展哪些字段


1. 背景
from_pretrained原子操作:
3依赖2的切分等处理,因此需要2输出
切分等(不需要输出sharding信息,这部分由param自己的placements, device_mesh属性来记录)操作的信息,并规范其格式2. 格式
from transformers.core_model_loading import ( WeightConverter, WeightRenaming, convert_and_load_state_dict_in_model, revert_weight_conversion, ) WeightMapping: list[WeightConverter | WeightRenaming] | None = None3. 流程
3.1. 伪代码
def from_pretrained(...): # 1. init model skeleton with device_meta_init_ctx: model = _init_model(...) # 1.5. init ModelConversionMetadata # - WeightMapping: from transformers # - NO DTensor Metadata: as param attr managed by FSDP2Manager weights_mapping = init_weights_mapping(model, ...) # 2. parallize & ... # 2.1. perf module perf_model, weights_mapping = apply_perf_model(model, weights_mapping) # 2.1. sharding plan sharded_model = apply_sharding_plan(model, ...) # 2.2. fsdp final_model = fsdp2_manager.parallelize(sharded_model, ...) # 3. load checkpoint # Params of final_model have attr: # - placements # - device_mesh # for checkpoint sharding & loading load_checkpoint(final_model, ckpt_file, weights_mapping)