(script_args, training_args, model_args)
| 18 | |
| 19 | # ----------------------- Main Script ----------------------- |
| 20 | def main(script_args, training_args, model_args): |
| 21 | reward_funcs = [reward_funcs_registry[func] for func in script_args.reward_funcs] |
| 22 | print("reward_funcs:", reward_funcs) |
| 23 | |
| 24 | |
| 25 | model_init_kwargs = {} |
| 26 | model_init_kwargs["attn_implementation"] = model_args.attn_implementation |
| 27 | model_init_kwargs["torch_dtype"] = torch.bfloat16 |
| 28 | model_init_kwargs["trust_remote_code"] = True |
| 29 | model_id = model_args.model_name_or_path |
| 30 | if "minicpm" in model_id.lower(): |
| 31 | if "minicpm-o" in model_id.lower(): |
| 32 | model_init_kwargs["init_tts"] = False |
| 33 | model_init_kwargs["init_audio"] = False |
| 34 | |
| 35 | model = AutoModelForCausalLM.from_pretrained(model_id, **model_init_kwargs) |
| 36 | processing_class = AutoProcessor.from_pretrained(model_id,trust_remote_code=True) |
| 37 | # processing_class.pad_token_id = processing_class.tokenizer.pad_token_id |
| 38 | # if processing_class.pad_token_id is None: |
| 39 | # processing_class.tokenizer.pad_token_id = 2 |
| 40 | # processing_class.pad_token_id = 2 |
| 41 | # processing_class.padding_side = "left" |
| 42 | # processing_class.tokenizer.padding_side = "left" |
| 43 | device_mesh = None |
| 44 | if training_args.tensor_parallel_size is not None: |
| 45 | tp_size = int(training_args.tensor_parallel_size) |
| 46 | world_size = dist.get_world_size() |
| 47 | if world_size % tp_size != 0: |
| 48 | raise ValueError( |
| 49 | f"world_size {world_size} must be divisible by tensor_parallel_size {tp_size}" |
| 50 | ) |
| 51 | dp_size = world_size // tp_size |
| 52 | |
| 53 | device_mesh = dist.device_mesh.DeviceMesh( |
| 54 | "cuda", |
| 55 | mesh=torch.arange(world_size).reshape((dp_size, tp_size)), |
| 56 | mesh_dim_names=("dp", "tp"), |
| 57 | ) |
| 58 | |
| 59 | tp_mesh = device_mesh["tp"] |
| 60 | dp_mesh = device_mesh["dp"] |
| 61 | |
| 62 | from torch.distributed.tensor.parallel import ColwiseParallel,RowwiseParallel,parallelize_module,SequenceParallel,PrepareModuleInput,PrepareModuleOutput |
| 63 | from torch.distributed.tensor import Replicate,Shard |
| 64 | |
| 65 | layer_tp_plan = { |
| 66 | "llm.model.layers.*.mlp": PrepareModuleInput( |
| 67 | input_layouts=(Shard(1),), |
| 68 | desired_input_layouts=(Replicate(),) |
| 69 | ), |
| 70 | "llm.model.layers.*.mlp.up_proj": ColwiseParallel(), |
| 71 | "llm.model.layers.*.mlp.gate_proj": ColwiseParallel(), |
| 72 | "llm.model.layers.*.mlp.down_proj": RowwiseParallel(output_layouts=Shard(1)), |
| 73 | "llm.model.layers.*.self_attn": PrepareModuleInput( |
| 74 | input_kwarg_layouts={ |
| 75 | "hidden_states": Shard(1), |
| 76 | "attention_mask": None, |
| 77 | "position_ids": None, |
no test coverage detected