MCPcopy Create free account
hub / github.com/OpenBMB/AgentCPM-GUI / main

Function main

rft/grpo.py:20–197  ·  view source on GitHub ↗
(script_args, training_args, model_args)

Source from the content-addressed store, hash-verified

18
19# ----------------------- Main Script -----------------------
20def 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,

Callers 1

grpo.pyFile · 0.70

Calls 3

save_modelMethod · 0.95
fsdp2_prepare_modelFunction · 0.90
AsyncRLGRPOTrainerClass · 0.90

Tested by

no test coverage detected