SIGN IN SIGN UP

🚨 TP dtensor API inference + training (#47579)

* add distributed config

* Add native FSDP2 module and migrate FSDP imports (Phase A PR-2).

Move FSDP2 wrapping and plan verification to distributed/fsdp.py, keep
integrations/fsdp.py as a backward-compatible re-export, and update core
call sites to import from transformers.distributed.fsdp.

* linting

* unecessary

* copyright edit

* revert

* add shard on read

* jsut shard on read

* cleaning

* linting

* fix

* fix

* remove redundant test file

* Update src/transformers/distributed/fsdp.py

naming

Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com>

* avoid looping, just look at dict

* expand_fsdp returns reshard_targets, no_reshard_targets right away

* better _resolve_tied_embed_lm_head_plan

* cleaning

* ruff

* more robust detection of embed and lm_head

* cleaning

* ruff

* typo

* cleaner

* cleaner

* typo

* refactor dense path + apply_contiguous_shard

* linting

* cleaning

* refactor _apply_strided_shard

* better

* refactor _slice_and_cat

* better comment

* refactor moe dtensor shard ops

* better comment

* comment

* cleaning

* linting

* Add FSDP orchestration: mesh init, distribute-before-load, and DCP save.

Wire distributed_config from_pretrained/save_pretrained alongside the legacy tp_plan path, add distributed/utils.py for mesh orchestration and checkpoint I/O, and extend sharding_utils with DTensor gather/optimizer fusion helpers needed by save/load.

* add fsdp plan to 2 models for now

* add tests fsdp mixin

* linting

* refactor test fsdp mixin

* test fsdp mixin cleaning

* remove fsdp policy in tests + trim down further

* test fsdp clean

* restore test_modeling_utils

* linting

* start trim down stuff

* fix

* breaking: cleaning modeling_utils.py

* load path with fsdp (dtensor) and tp (old tp) is linked

* linting

* add saving

* styling

* fix tp ci

* add fsdp to ci

* linting

* pick one model only for this PR

* restore

* trigger fsdp ci

* doc cleaning +  tp_size remove

* fix tp ci for ep

* edit doc

* move distributed function to utils + guarding

* linting

* expand_fsdp_plan iterate over modules

* comment about tie embedding

* add comment tied embedding

* add DistributedMixin

* some cleaning

* cleaning + comment

* rename function for clarity

* Apply suggestion from @ArthurZucker

Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com>

* doc

* comment

* linting

* refactor

* abstract to mixin

* typo

* Add FSDP plans to all models from distributed branch.

Port base_model_fsdp_plan and ForCausalLM _fsdp_plan entries from PR #46269
and expand FSDP distributed test coverage to the pilot model subset.

* fsdp plans

* linting

* linting

* Add distributed runtime utils and DistributedMixin (FSDP orchestration 1/3).

Introduce distributed/utils.py and DistributedMixin, defer DistributedConfig
validation to load time, and refactor PreTrainedModel plan properties without
changing the from_pretrained distributed_config API yet.

* Wire DistributedConfig through from_pretrained and save_pretrained (FSDP orchestration 2/3).

Route distributed loading and saving through DistributedMixin, migrate TP tests
and docs off tp_plan="auto", and add FSDP gather/DCP save paths.

* Add FSDP CI and end-to-end FSDP tests (FSDP orchestration 3/3).

Add FSDPTesterMixin, cohere2_moe base_fsdp_plan, dedicated fsdp_ci job, and
pytest markers for distributed FSDP load/save/generation coverage.

* addd ep_plan

* restore validate module

* Wire DistributedConfig through from_pretrained and save_pretrained.

Route distributed load/save orchestration through DistributedMixin so TP and FSDP paths share the same entry points.

* revert

* inline distribute_model

* revert

* remove saving/loading

* leaner mixin

* downgrade torch version guarding

* remove

* linting

* revert

* revert

* post_init() parallel plan move to mixin

* revert tp mixin

* add save/load

* only FSDP save/load for now

* revert

* refactor

* modular

* ea

* begin migration TP

* clean up

* add it to pretrained model

* linting

* fix the test by moving in init class the fsdp plan instead of post init

* edit

* replace everything

* begin migration TP

* clean up

* revert merge conflict

* local params for forward

* cleaning

* migration from integration.tensor_parallel to   distributed.tensor_parallel

* revert models

* revert

* fix

* breaking: cleaner way to_local for quantize weight(almost done need to fix backward)

* fix all tests

* TP dtensor handle natively the to_local() for deepgem + fp8 (#47634)

* Refactor distributed tensor_parallel module and TP mixin tests

* better approach to to_local

* pass as class method

* cleaning PackedColwise

* remove the use of to_local in kernels to offload this task to TensorParallel

* rowwise bias after redistribute

* make it easier to understand

* requires grad only for floating point numbers

* fix tp meagamoe fp8 with dtensor

* linting

* revert gitignore

* fix run slow quantization ci

* fix deepseek v4 ep backward tests

* fix deepseek glm4 moe tp backward pass

* Refactor tensor parallel loading logic by removing unused tp_plan handling and updating is_dtensor type hint for better type safety.

* linting

* bench: dtensor vs legacy TP

* bench: dtensor vs legacy TP (#47728)

* bench: dtensor vs legacy TP

* clean rowise

* refactor colwise

* comment

* comment

* refactor MoEExperParallel

* ruff

* add rowwise input

* remove helper script

* remove dead code in mxfp4

* remove old TP

* cleaning

* cleaning

* claening

* add test_shard_tensor_shape_consistency

* cleaning

* small fix

* typo

* comment

* rename colwise_gather_output to colwise_rep

* keep tp import to avoid breaking chnges

* remove async_op=True for redistribute

* linting

* renaming MoeExpertsParallel class

* ep router doc

* mlinter

* add todo

* fix

* remove unused keys from global_mapping to avoid BC

---------

Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com>
F
Ferdinand Mom committed
861f4c41ee46974debe5dde47dd0339146f4e6d2
Parent: d4c297d
Committed by GitHub <noreply@github.com> on 8/21/2026, 12:25:37 PM