MCPcopy Create free account
hub / github.com/PaddlePaddle/Paddle / wrapped

Function wrapped

python/paddle/distributed/auto_parallel/local_map.py:125–252  ·  view source on GitHub ↗
(process_mesh: ProcessMesh | None, *args, **kwargs)

Source from the content-addressed store, hash-verified

123 """
124
125 def wrapped(process_mesh: ProcessMesh | None, *args, **kwargs):
126 # Process input arguments
127 flat_dist_args = flatten(args)
128 if in_placements is not None:
129 assert len(in_placements) == len(flat_dist_args), (
130 f"in_placements length {len(in_placements)} does not match "
131 f"number of input args {len(flat_dist_args)}!"
132 )
133
134 flat_local_args = []
135 seen_dist_tensor = False
136
137 for idx, arg in enumerate(flat_dist_args):
138 if dist.auto_parallel.api.is_dist_tensor(arg):
139 dist_tensor = arg
140 if process_mesh is None:
141 if paddle.in_dynamic_mode():
142 process_mesh = dist_tensor.process_mesh
143 else:
144 process_mesh = dist_tensor.dist_attr().process_mesh
145
146 seen_dist_tensor = True
147
148 if in_placements is not None:
149 in_placement = in_placements[idx]
150 if in_placement is None:
151 if paddle.in_dynamic_mode():
152 in_placement = dist_tensor.placements
153 else:
154 in_placement = dist_tensor.dist_attr().placements
155 else:
156 if paddle.in_dynamic_mode():
157 if in_placement != dist_tensor.placements:
158 if reshard_inputs:
159 dist_tensor = dist.reshard(
160 dist_tensor, process_mesh, in_placement
161 )
162 else:
163 raise ValueError(
164 f"in_placement {in_placement} does not match dist_tensor.placements {dist_tensor.placements}"
165 )
166
167 else:
168 if (
169 in_placement
170 != dist_tensor.dist_attr().placements
171 ):
172 if reshard_inputs:
173 dist_tensor = dist.reshard(
174 dist_tensor, process_mesh, in_placement
175 )
176 else:
177 raise ValueError(
178 f"in_placement {in_placement} does not match dist_tensor.dist_attr().placements {dist_tensor.dist_attr().placements}"
179 "If reshard_inputs is wanted, set "
180 "reshard_inputs=True to local_map."
181 )
182 local_tensor = dist.auto_parallel.api.dtensor_to_local(

Callers

nothing calls this directly

Calls 9

flattenFunction · 0.90
pack_sequence_asFunction · 0.90
ValueErrorClass · 0.85
is_dist_tensorMethod · 0.80
funcFunction · 0.50
typeFunction · 0.50
dist_attrMethod · 0.45
reshardMethod · 0.45
appendMethod · 0.45

Tested by

no test coverage detected