(process_mesh: ProcessMesh | None, *args, **kwargs)
| 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( |
nothing calls this directly
no test coverage detected