| 331 | |
| 332 | |
| 333 | def create_parameter( |
| 334 | dtype, |
| 335 | shape, |
| 336 | name=None, |
| 337 | **kwargs, |
| 338 | ): |
| 339 | if 'initializer' not in kwargs: |
| 340 | raise ValueError( |
| 341 | "initializer is None, if you want to create parameter, please pass its initializer." |
| 342 | ) |
| 343 | if dtype is not None: |
| 344 | if not isinstance(dtype, DataType): |
| 345 | dtype = convert_np_dtype_to_dtype_(dtype) |
| 346 | value_name = name |
| 347 | if not value_name: |
| 348 | value_name = unique_name.generate('parameter') |
| 349 | startup_program = default_startup_program() |
| 350 | main_program = default_main_program() |
| 351 | parameter_meta = ParameterMeta(shape, dtype) |
| 352 | |
| 353 | is_dist = False |
| 354 | if ( |
| 355 | 'placements' in kwargs |
| 356 | and kwargs['placements'] |
| 357 | and 'process_mesh' in kwargs |
| 358 | and kwargs['process_mesh'] |
| 359 | ): |
| 360 | is_dist = True |
| 361 | |
| 362 | def to_dist(value): |
| 363 | import paddle |
| 364 | import paddle.distributed as dist |
| 365 | |
| 366 | process_mesh = kwargs['process_mesh'] |
| 367 | dim_map, partial_status = dist.auto_parallel.placement_type.to_dim_map( |
| 368 | kwargs['placements'], len(shape) |
| 369 | ) |
| 370 | dist_attr = paddle.base.libpaddle.pir.create_tensor_dist_attribute( |
| 371 | process_mesh, dim_map, partial_status |
| 372 | ) |
| 373 | dist_type = paddle.base.libpaddle.pir.cvt_to_dist_type( |
| 374 | value.type(), dist_attr |
| 375 | ) |
| 376 | value.set_type(dist_type) |
| 377 | op_dist_attr = paddle.base.libpaddle.pir.create_op_dist_attribute( |
| 378 | process_mesh, [], [dist_attr] |
| 379 | ) |
| 380 | value.get_defining_op().dist_attr = op_dist_attr |
| 381 | |
| 382 | with program_guard(startup_program): |
| 383 | initializer = kwargs['initializer'] |
| 384 | init_result = initializer( |
| 385 | parameter_meta, startup_program.global_block() |
| 386 | ) |
| 387 | init_result.persistable = True |
| 388 | if is_dist: |
| 389 | to_dist(init_result) |
| 390 | |