Skip to content

Commit c649d1d

Browse files
authored
[Auto-parallel] Clean dy_auto sharding opt flag (#73967)
1 parent 5b4021a commit c649d1d

4 files changed

Lines changed: 25 additions & 48 deletions

File tree

python/paddle/amp/auto_cast.py

Lines changed: 10 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -693,22 +693,13 @@ def master_grad_hook():
693693
].append(param)
694694
amp_global_state().already_classify_params_meshes = True
695695

696-
if os.getenv("FLAGS_enable_tensor_fusion") not in [
697-
"True",
698-
"true",
699-
"1",
700-
] and os.getenv("FLAGS_enable_main_grad") not in [
701-
"True",
702-
"true",
703-
"1",
704-
]:
705-
if len(amp_global_state().mesh2params):
706-
for _, params in amp_global_state().mesh2params.items():
707-
core.eager.set_master_grads(params)
708-
else:
709-
core.eager.set_master_grads(
710-
amp_global_state().model_parameters
711-
)
696+
if len(amp_global_state().mesh2params):
697+
for _, params in amp_global_state().mesh2params.items():
698+
core.eager.set_master_grads(params)
699+
else:
700+
core.eager.set_master_grads(
701+
amp_global_state().model_parameters
702+
)
712703

713704
amp_global_state().already_register_final_backward_hook = False
714705

@@ -750,8 +741,9 @@ def param_hook(tmp_grad):
750741
if not hasattr(param, "main_grad"):
751742
param.main_grad = None
752743
param._register_grad_hook(_update_main_grad_hook(param))
753-
754-
core.eager._add_backward_final_hook(master_grad_hook)
744+
os.environ["FLAGS_enable_tensor_fusion"] = "0"
745+
else:
746+
core.eager._add_backward_final_hook(master_grad_hook)
755747
amp_global_state().already_register_final_backward_hook = True
756748

757749
if tracer:

python/paddle/distributed/auto_parallel/api.py

Lines changed: 12 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -1163,14 +1163,8 @@ def __init__(self, optimizer, shard_fn=None, gradient_accumulation_steps=1):
11631163
self._mp_group = None
11641164
self.do_tensor_fusion_once = True
11651165
self._strategy = Strategy()
1166-
self.enable_tensor_fusion = os.getenv("FLAGS_enable_tensor_fusion") in [
1167-
"True",
1168-
"true",
1169-
"1",
1170-
]
1171-
self.enable_sharding_overlap = os.getenv(
1172-
"FLAGS_enable_sharding_overlap"
1173-
) in ["True", "true", "1"]
1166+
self.enable_tensor_fusion = False
1167+
self.enable_sharding_overlap = False
11741168

11751169
def _set_and_check_sharding_prop_from_param(self):
11761170
global_mesh = fleet.auto.get_mesh()
@@ -1507,14 +1501,14 @@ def get_mesh(pp_idx=0):
15071501
self.param_storage[idx].is_sync = False
15081502

15091503
def _enable_tensor_fusion(self):
1510-
# TODO: enable after clear FLAGS_enable_tensor_fusion
1511-
# self.enable_tensor_fusion = True
1512-
pass
1504+
os.environ["FLAGS_enable_tensor_fusion"] = "1"
1505+
self.enable_tensor_fusion = True
1506+
self._shard_fn._enable_tensor_fusion()
15131507

15141508
def _enable_sharding_overlap(self, layers):
15151509
if hasattr(layers, 'config') and layers.config.get("to_static", False):
15161510
return
1517-
# self.enable_sharding_overlap = True
1511+
self.enable_sharding_overlap = True
15181512
if not isinstance(layers, paddle.nn.Layer):
15191513
raise RuntimeError(
15201514
f"`layers` must be `paddle.nn.Layer` but got {type(layers)}"
@@ -1971,15 +1965,19 @@ def __init__(self, mesh, sharding_mesh_dim):
19711965
self._mesh = mesh
19721966
self._sharding_axis = 0
19731967
self._sharding_mesh_dim = sharding_mesh_dim
1968+
self.enable_tensor_fusion = False
19741969

19751970
def _set_sharding_axis(self, sharding_axis):
19761971
self._sharding_axis = sharding_axis
19771972

1973+
def _enable_tensor_fusion(self):
1974+
self.enable_tensor_fusion = True
1975+
19781976
def shard_master_weight(
19791977
self, param: Tensor, master_weight: Tensor
19801978
) -> Tensor:
19811979
if param.is_dist():
1982-
if os.getenv("FLAGS_enable_tensor_fusion") in ["True", "true", "1"]:
1980+
if self.enable_tensor_fusion:
19831981
placements = param.placements
19841982
else:
19851983
placements = get_placement_with_sharding(
@@ -2115,10 +2113,7 @@ def __call__(self, key: str, param: Tensor, tensor: Tensor) -> Tensor:
21152113
return tensor
21162114

21172115
# Only deal with momentum in optimizer, beta should be replicated cross param's mesh
2118-
if (
2119-
os.getenv("FLAGS_enable_tensor_fusion") not in ["True", "true", "1"]
2120-
and 'beta' not in key
2121-
):
2116+
if not self.enable_tensor_fusion and 'beta' not in key:
21222117
placements = get_placement_with_sharding(param, self._sharding_axis)
21232118
else:
21242119
placements = [

python/paddle/optimizer/optimizer.py

Lines changed: 3 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -2016,11 +2016,9 @@ def step(self) -> None:
20162016
for param in self._param_groups:
20172017
if param.stop_gradient:
20182018
continue
2019-
if os.getenv("FLAGS_enable_tensor_fusion") in [
2020-
"True",
2021-
"true",
2022-
"1",
2023-
] or os.getenv("FLAGS_enable_main_grad") in [
2019+
if getattr(self, 'enable_tensor_fusion', False) or os.getenv(
2020+
"FLAGS_enable_main_grad"
2021+
) in [
20242022
"True",
20252023
"true",
20262024
"1",

test/auto_parallel/semi_auto_parallel_sharding_stage_1.py

Lines changed: 0 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -141,9 +141,6 @@ def test_sharding_stage_1_overlap_to_static(self):
141141

142142
def test_pure_sharding_multi_mesh_stage_1_with_tensor_fusion(self):
143143
def run_sharding_test(enable_tensor_fusion):
144-
os.environ['FLAGS_enable_tensor_fusion'] = (
145-
'1' if enable_tensor_fusion else '0'
146-
)
147144
paddle.distributed.auto_parallel.set_mesh(self._multi_dim_mesh)
148145
self.set_random_seed()
149146
model = paddle.nn.Linear(10, 10)
@@ -176,7 +173,6 @@ def test_pure_sharding_multi_mesh_stage_1_with_tensor_fusion_with_chip(
176173
self,
177174
):
178175
dist.init_parallel_env()
179-
os.environ['FLAGS_enable_tensor_fusion'] = '1'
180176
paddle.distributed.auto_parallel.set_mesh(self._multi_dim_mesh)
181177
self.set_random_seed()
182178
model = paddle.nn.Linear(10, 10)
@@ -204,10 +200,6 @@ def test_pure_sharding_multi_mesh_stage_1_with_tensor_fusion_with_chip(
204200

205201
def test_pure_sharding_multi_mesh_stage_1_with_sharding_overlap(self):
206202
def run_sharding_test(enable_sharding_overlap):
207-
os.environ['FLAGS_enable_tensor_fusion'] = '1'
208-
os.environ['FLAGS_enable_sharding_overlap'] = (
209-
'1' if enable_sharding_overlap else '0'
210-
)
211203
paddle.distributed.auto_parallel.set_mesh(self._multi_dim_mesh)
212204
self.set_random_seed()
213205
model = paddle.nn.Linear(10, 10)

0 commit comments

Comments
 (0)